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,270 @@
1
+ """Graph topology construction strategies from tabular data."""
2
+
3
+ from typing import Any, Dict, List, Optional, Sequence, Tuple, Union
4
+ import numpy as np
5
+
6
+
7
+ class KNNGraphBuilder:
8
+ r"""Constructs a k-nearest-neighbors graph from a node feature matrix.
9
+
10
+ Args:
11
+ k: Number of nearest neighbors per node. (default: ``5``)
12
+ metric: Distance metric (``"cosine"``, ``"euclidean"``, or ``"manhattan"``).
13
+ (default: ``"cosine"``)
14
+ loop: Whether to include self-loops. (default: ``False``)
15
+ bidirectional: Whether to make the resulting graph undirected. (default: ``True``)
16
+ """
17
+
18
+ def __init__(
19
+ self,
20
+ k: int = 5,
21
+ metric: str = "cosine",
22
+ loop: bool = False,
23
+ bidirectional: bool = True,
24
+ ):
25
+ self.k = k
26
+ self.metric = metric.lower()
27
+ self.loop = loop
28
+ self.bidirectional = bidirectional
29
+
30
+ def __call__(self, x: np.ndarray, **kwargs) -> Tuple[np.ndarray, Optional[np.ndarray]]:
31
+ num_nodes = x.shape[0]
32
+ if num_nodes == 0:
33
+ return np.empty((2, 0), dtype=np.int64), np.empty((0, 1), dtype=np.float32)
34
+
35
+ actual_k = min(self.k if self.loop else self.k + 1, num_nodes)
36
+
37
+ if self.metric == "cosine":
38
+ norms = np.linalg.norm(x, axis=1, keepdims=True)
39
+ norms[norms < 1e-8] = 1.0
40
+ x_norm = x / norms
41
+ sim_matrix = np.dot(x_norm, x_norm.T)
42
+ # Higher similarity is closer
43
+ dist_matrix = 1.0 - sim_matrix
44
+ elif self.metric == "euclidean":
45
+ diff = x[:, np.newaxis, :] - x[np.newaxis, :, :]
46
+ dist_matrix = np.sqrt(np.sum(diff ** 2, axis=-1))
47
+ elif self.metric == "manhattan":
48
+ diff = x[:, np.newaxis, :] - x[np.newaxis, :, :]
49
+ dist_matrix = np.sum(np.abs(diff), axis=-1)
50
+ else:
51
+ raise ValueError(f"Unknown metric '{self.metric}'. Supported: 'cosine', 'euclidean', 'manhattan'.")
52
+
53
+ src_list = []
54
+ dst_list = []
55
+ dist_list = []
56
+
57
+ for i in range(num_nodes):
58
+ row_dists = dist_matrix[i]
59
+ nearest_indices = np.argsort(row_dists)
60
+ count = 0
61
+ for neighbor in nearest_indices:
62
+ if not self.loop and neighbor == i:
63
+ continue
64
+ src_list.append(i)
65
+ dst_list.append(neighbor)
66
+ dist_list.append(row_dists[neighbor])
67
+ count += 1
68
+ if count >= self.k:
69
+ break
70
+
71
+ if self.bidirectional:
72
+ # Add reverse edges if not already present
73
+ edge_set = set(zip(src_list, dst_list))
74
+ for s, d, w in list(zip(src_list, dst_list, dist_list)):
75
+ if (d, s) not in edge_set:
76
+ src_list.append(d)
77
+ dst_list.append(s)
78
+ dist_list.append(w)
79
+ edge_set.add((d, s))
80
+
81
+ edge_index = np.array([src_list, dst_list], dtype=np.int64)
82
+ edge_attr = np.array(dist_list, dtype=np.float32).reshape(-1, 1) if dist_list else np.empty((0, 1), dtype=np.float32)
83
+ return edge_index, edge_attr
84
+
85
+
86
+ class SimilarityGraphBuilder:
87
+ r"""Constructs a graph connecting node pairs whose pairwise similarity exceeds a threshold.
88
+
89
+ Args:
90
+ threshold: Minimum similarity required to create an edge. (default: ``0.7``)
91
+ metric: Similarity function (``"cosine"`` or ``"rbf"``). (default: ``"cosine"``)
92
+ gamma: Bandwidth parameter for RBF kernel. (default: ``1.0``)
93
+ loop: Whether to include self-loops. (default: ``False``)
94
+ """
95
+
96
+ def __init__(
97
+ self,
98
+ threshold: float = 0.7,
99
+ metric: str = "cosine",
100
+ gamma: float = 1.0,
101
+ loop: bool = False,
102
+ ):
103
+ self.threshold = threshold
104
+ self.metric = metric.lower()
105
+ self.gamma = gamma
106
+ self.loop = loop
107
+
108
+ def __call__(self, x: np.ndarray, **kwargs) -> Tuple[np.ndarray, Optional[np.ndarray]]:
109
+ num_nodes = x.shape[0]
110
+ if num_nodes == 0:
111
+ return np.empty((2, 0), dtype=np.int64), np.empty((0, 1), dtype=np.float32)
112
+
113
+ if self.metric == "cosine":
114
+ norms = np.linalg.norm(x, axis=1, keepdims=True)
115
+ norms[norms < 1e-8] = 1.0
116
+ x_norm = x / norms
117
+ sim_matrix = np.dot(x_norm, x_norm.T)
118
+ elif self.metric == "rbf":
119
+ diff = x[:, np.newaxis, :] - x[np.newaxis, :, :]
120
+ sq_dist = np.sum(diff ** 2, axis=-1)
121
+ sim_matrix = np.exp(-self.gamma * sq_dist)
122
+ else:
123
+ raise ValueError(f"Unknown metric '{self.metric}'. Supported: 'cosine', 'rbf'.")
124
+
125
+ if not self.loop:
126
+ np.fill_diagonal(sim_matrix, -np.inf)
127
+
128
+ src, dst = np.where(sim_matrix >= self.threshold)
129
+ weights = sim_matrix[src, dst].astype(np.float32).reshape(-1, 1)
130
+ edge_index = np.stack([src, dst], axis=0).astype(np.int64)
131
+ return edge_index, weights
132
+
133
+
134
+ class SharedEntityGraphBuilder:
135
+ r"""Connects tabular rows that share one or more categorical identifier values.
136
+ (e.g., users sharing the same IP address, device, category, or cluster).
137
+
138
+ Args:
139
+ entity_cols: List of column names to check for shared values.
140
+ max_degree: Maximum number of neighbors created per shared entity (to avoid supernode explosion).
141
+ (default: ``50``)
142
+ loop: Whether to include self-loops. (default: ``False``)
143
+ """
144
+
145
+ def __init__(
146
+ self,
147
+ entity_cols: Sequence[str],
148
+ max_degree: int = 50,
149
+ loop: bool = False,
150
+ ):
151
+ self.entity_cols = list(entity_cols)
152
+ self.max_degree = max_degree
153
+ self.loop = loop
154
+
155
+ def __call__(self, x: np.ndarray, df_or_dict: Optional[Any] = None, **kwargs) -> Tuple[np.ndarray, Optional[np.ndarray]]:
156
+ if df_or_dict is None:
157
+ raise ValueError("SharedEntityGraphBuilder requires 'df_or_dict' containing the entity columns.")
158
+
159
+ num_nodes = x.shape[0]
160
+ src_list = []
161
+ dst_list = []
162
+
163
+ from k3_node.etl.encoders import _get_column_values
164
+
165
+ for col in self.entity_cols:
166
+ vals = _get_column_values(df_or_dict, col)
167
+ val_to_rows: Dict[Any, List[int]] = {}
168
+ for row_idx, val in enumerate(vals):
169
+ if val is None or (isinstance(val, float) and np.isnan(val)) or val == "":
170
+ continue
171
+ val_to_rows.setdefault(val, []).append(row_idx)
172
+
173
+ for val, rows in val_to_rows.items():
174
+ if len(rows) > self.max_degree:
175
+ # Subsample if group is too large
176
+ sampled_rows = np.random.choice(rows, size=self.max_degree, replace=False).tolist()
177
+ else:
178
+ sampled_rows = rows
179
+
180
+ for i in sampled_rows:
181
+ for j in sampled_rows:
182
+ if not self.loop and i == j:
183
+ continue
184
+ src_list.append(i)
185
+ dst_list.append(j)
186
+
187
+ if not src_list:
188
+ return np.empty((2, 0), dtype=np.int64), np.empty((0, 1), dtype=np.float32)
189
+
190
+ edges = list(set(zip(src_list, dst_list)))
191
+ src_arr = np.array([e[0] for e in edges], dtype=np.int64)
192
+ dst_arr = np.array([e[1] for e in edges], dtype=np.int64)
193
+ edge_index = np.stack([src_arr, dst_arr], axis=0)
194
+ edge_attr = np.ones((len(edges), 1), dtype=np.float32)
195
+ return edge_index, edge_attr
196
+
197
+
198
+ class SequentialGraphBuilder:
199
+ r"""Connects tabular rows sequentially in order of an index or timestamp column,
200
+ optionally partitioned by a group entity column.
201
+
202
+ Args:
203
+ order_col: Optional column name used to sort rows (e.g. timestamp or sequence index).
204
+ group_by_col: Optional column name to partition sequences (e.g. user_id or session_id).
205
+ window_size: Number of forward/backward sequential steps to connect. (default: ``1``)
206
+ bidirectional: Whether to create undirected edges. (default: ``True``)
207
+ """
208
+
209
+ def __init__(
210
+ self,
211
+ order_col: Optional[str] = None,
212
+ group_by_col: Optional[str] = None,
213
+ window_size: int = 1,
214
+ bidirectional: bool = True,
215
+ ):
216
+ self.order_col = order_col
217
+ self.group_by_col = group_by_col
218
+ self.window_size = window_size
219
+ self.bidirectional = bidirectional
220
+
221
+ def __call__(self, x: np.ndarray, df_or_dict: Optional[Any] = None, **kwargs) -> Tuple[np.ndarray, Optional[np.ndarray]]:
222
+ num_nodes = x.shape[0]
223
+ if num_nodes == 0:
224
+ return np.empty((2, 0), dtype=np.int64), np.empty((0, 1), dtype=np.float32)
225
+
226
+ from k3_node.etl.encoders import _get_column_values
227
+
228
+ if self.group_by_col is not None and df_or_dict is not None:
229
+ groups = _get_column_values(df_or_dict, self.group_by_col)
230
+ group_to_indices: Dict[Any, List[int]] = {}
231
+ for idx, g in enumerate(groups):
232
+ group_to_indices.setdefault(g, []).append(idx)
233
+ else:
234
+ group_to_indices = {"all": list(range(num_nodes))}
235
+
236
+ if self.order_col is not None and df_or_dict is not None:
237
+ order_vals = _get_column_values(df_or_dict, self.order_col)
238
+ else:
239
+ order_vals = None
240
+
241
+ src_list = []
242
+ dst_list = []
243
+
244
+ for group_name, row_indices in group_to_indices.items():
245
+ if order_vals is not None:
246
+ sorted_indices = sorted(row_indices, key=lambda idx: order_vals[idx])
247
+ else:
248
+ sorted_indices = row_indices
249
+
250
+ n_seq = len(sorted_indices)
251
+ for i in range(n_seq):
252
+ curr_node = sorted_indices[i]
253
+ for step in range(1, self.window_size + 1):
254
+ if i + step < n_seq:
255
+ next_node = sorted_indices[i + step]
256
+ src_list.append(curr_node)
257
+ dst_list.append(next_node)
258
+ if self.bidirectional:
259
+ src_list.append(next_node)
260
+ dst_list.append(curr_node)
261
+
262
+ if not src_list:
263
+ return np.empty((2, 0), dtype=np.int64), np.empty((0, 1), dtype=np.float32)
264
+
265
+ edges = list(set(zip(src_list, dst_list)))
266
+ src_arr = np.array([e[0] for e in edges], dtype=np.int64)
267
+ dst_arr = np.array([e[1] for e in edges], dtype=np.int64)
268
+ edge_index = np.stack([src_arr, dst_arr], axis=0)
269
+ edge_attr = np.ones((len(edges), 1), dtype=np.float32)
270
+ return edge_index, edge_attr
@@ -0,0 +1,201 @@
1
+ """Relational (multi-table) to Heterogeneous Graph ETL converter."""
2
+
3
+ from typing import Any, Dict, List, Optional, Sequence, Tuple, Union
4
+ import numpy as np
5
+
6
+ from k3_node.data import HeteroData, Data
7
+ from k3_node.etl.encoders import TabularEncoder, _get_column_names, _get_column_values
8
+
9
+
10
+ NodeType = str
11
+ EdgeType = Tuple[str, str, str]
12
+
13
+
14
+ class RelationalToGraph:
15
+ r"""ETL pipeline converting multi-table relational databases / DataFrames into
16
+ a heterogeneous graph :class:`k3_node.data.HeteroData` object.
17
+
18
+ Args:
19
+ id_cols (dict): Mapping from :obj:`NodeType` to the primary key column name.
20
+ (e.g., ``{"user": "user_id", "movie": "movie_id"}``).
21
+ edge_cols (dict): Mapping from :obj:`EdgeType` to the tuple of foreign key column names
22
+ ``(source_id_col, target_id_col)``.
23
+ (e.g., ``{("user", "rates", "movie"): ("user_id", "movie_id")}``).
24
+ feature_cols (dict, optional): Mapping from :obj:`NodeType` to a list of feature columns.
25
+ If omitted, all columns except the primary key are encoded.
26
+ edge_attr_cols (dict, optional): Mapping from :obj:`EdgeType` to a list of edge attribute columns.
27
+ node_target_cols (dict, optional): Mapping from :obj:`NodeType` to the target label column.
28
+ edge_target_cols (dict, optional): Mapping from :obj:`EdgeType` to the target edge label column.
29
+ """
30
+
31
+ def __init__(
32
+ self,
33
+ id_cols: Dict[NodeType, str],
34
+ edge_cols: Dict[EdgeType, Tuple[str, str]],
35
+ feature_cols: Optional[Dict[NodeType, List[str]]] = None,
36
+ edge_attr_cols: Optional[Dict[EdgeType, List[str]]] = None,
37
+ node_target_cols: Optional[Dict[NodeType, str]] = None,
38
+ edge_target_cols: Optional[Dict[EdgeType, str]] = None,
39
+ ):
40
+ self.id_cols = id_cols
41
+ self.edge_cols = edge_cols
42
+ self.feature_cols = feature_cols or {}
43
+ self.edge_attr_cols = edge_attr_cols or {}
44
+ self.node_target_cols = node_target_cols or {}
45
+ self.edge_target_cols = edge_target_cols or {}
46
+
47
+ self.node_encoders_: Dict[NodeType, TabularEncoder] = {}
48
+ self.edge_encoders_: Dict[EdgeType, TabularEncoder] = {}
49
+ self.id_maps_: Dict[NodeType, Dict[Any, int]] = {}
50
+ self.inverse_id_maps_: Dict[NodeType, Dict[int, Any]] = {}
51
+
52
+ def fit(self, nodes: Dict[NodeType, Any], edges: Optional[Dict[EdgeType, Any]] = None):
53
+ r"""Fits encoders and builds entity ID mappings across all tables."""
54
+ # 1. Map node IDs and fit node feature encoders
55
+ for node_type, table in nodes.items():
56
+ id_col = self.id_cols[node_type]
57
+ raw_ids = _get_column_values(table, id_col)
58
+
59
+ # Unique contiguous ID mapping
60
+ unique_ids = []
61
+ seen = set()
62
+ for rid in raw_ids:
63
+ if rid not in seen:
64
+ seen.add(rid)
65
+ unique_ids.append(rid)
66
+
67
+ id_map = {rid: i for i, rid in enumerate(unique_ids)}
68
+ inv_map = {i: rid for i, rid in enumerate(unique_ids)}
69
+ self.id_maps_[node_type] = id_map
70
+ self.inverse_id_maps_[node_type] = inv_map
71
+
72
+ # Feature columns
73
+ all_cols = _get_column_names(table)
74
+ target_col = self.node_target_cols.get(node_type)
75
+ ignore_cols = {id_col}
76
+ if target_col:
77
+ ignore_cols.add(target_col)
78
+
79
+ if node_type in self.feature_cols:
80
+ feat_cols = [c for c in self.feature_cols[node_type] if c in all_cols and c not in ignore_cols]
81
+ else:
82
+ feat_cols = [c for c in all_cols if c not in ignore_cols]
83
+
84
+ encoder = TabularEncoder()
85
+ if feat_cols:
86
+ encoder.fit(table, columns=feat_cols)
87
+ self.node_encoders_[node_type] = encoder
88
+
89
+ # 2. Fit edge encoders if edge attributes are specified
90
+ if edges:
91
+ for edge_type, table in edges.items():
92
+ if edge_type in self.edge_attr_cols:
93
+ attr_cols = self.edge_attr_cols[edge_type]
94
+ edge_enc = TabularEncoder()
95
+ edge_enc.fit(table, columns=attr_cols)
96
+ self.edge_encoders_[edge_type] = edge_enc
97
+
98
+ return self
99
+
100
+ def transform(
101
+ self,
102
+ nodes: Dict[NodeType, Any],
103
+ edges: Optional[Dict[EdgeType, Any]] = None,
104
+ ) -> HeteroData:
105
+ r"""Constructs a :class:`k3_node.data.HeteroData` instance from relational tables."""
106
+ hetero_data = HeteroData()
107
+
108
+ # 1. Process node tables
109
+ for node_type, table in nodes.items():
110
+ encoder = self.node_encoders_[node_type]
111
+ if encoder.column_order_:
112
+ x = encoder.transform(table)
113
+ hetero_data[node_type].x = x
114
+ else:
115
+ num_nodes = len(self.id_maps_[node_type])
116
+ hetero_data[node_type].num_nodes = num_nodes
117
+
118
+ # Target labels y
119
+ if node_type in self.node_target_cols:
120
+ target_col = self.node_target_cols[node_type]
121
+ raw_y = _get_column_values(table, target_col)
122
+ if all(isinstance(v, (int, np.integer)) for v in raw_y if v is not None):
123
+ hetero_data[node_type].y = np.array(raw_y, dtype=np.int64)
124
+ elif all(isinstance(v, (float, int, np.floating, np.integer)) for v in raw_y if v is not None):
125
+ hetero_data[node_type].y = np.array(raw_y, dtype=np.float32)
126
+ else:
127
+ unique_c = sorted(list(set(raw_y)))
128
+ mapping = {c: i for i, c in enumerate(unique_c)}
129
+ hetero_data[node_type].y = np.array([mapping[c] for c in raw_y], dtype=np.int64)
130
+
131
+ # 2. Process edge tables
132
+ if edges:
133
+ for edge_type, table in edges.items():
134
+ src_type, rel_name, dst_type = edge_type
135
+ src_col, dst_col = self.edge_cols[edge_type]
136
+
137
+ raw_srcs = _get_column_values(table, src_col)
138
+ raw_dsts = _get_column_values(table, dst_col)
139
+
140
+ src_map = self.id_maps_[src_type]
141
+ dst_map = self.id_maps_[dst_type]
142
+
143
+ valid_src = []
144
+ valid_dst = []
145
+ valid_indices = []
146
+
147
+ for row_idx, (s, d) in enumerate(zip(raw_srcs, raw_dsts)):
148
+ if s in src_map and d in dst_map:
149
+ valid_src.append(src_map[s])
150
+ valid_dst.append(dst_map[d])
151
+ valid_indices.append(row_idx)
152
+
153
+ edge_index = np.array([valid_src, valid_dst], dtype=np.int64)
154
+ hetero_data[edge_type].edge_index = edge_index
155
+
156
+ # Edge attributes
157
+ if edge_type in self.edge_encoders_:
158
+ edge_enc = self.edge_encoders_[edge_type]
159
+ full_attrs = edge_enc.transform(table)
160
+ hetero_data[edge_type].edge_attr = full_attrs[valid_indices]
161
+
162
+ # Edge target y
163
+ if edge_type in self.edge_target_cols:
164
+ t_col = self.edge_target_cols[edge_type]
165
+ raw_ey = _get_column_values(table, t_col)
166
+ filtered_ey = [raw_ey[idx] for idx in valid_indices]
167
+ hetero_data[edge_type].edge_label = np.array(filtered_ey, dtype=np.float32)
168
+
169
+ # Attach metadata
170
+ hetero_data.id_maps = self.id_maps_
171
+ hetero_data.inverse_id_maps = self.inverse_id_maps_
172
+ return hetero_data
173
+
174
+ def fit_transform(
175
+ self,
176
+ nodes: Dict[NodeType, Any],
177
+ edges: Optional[Dict[EdgeType, Any]] = None,
178
+ ) -> HeteroData:
179
+ r"""Fits encoders and builds the heterogeneous graph in a single call."""
180
+ return self.fit(nodes, edges).transform(nodes, edges)
181
+
182
+
183
+ def relational_to_graph(
184
+ nodes: Dict[NodeType, Any],
185
+ edges: Optional[Dict[EdgeType, Any]] = None,
186
+ id_cols: Optional[Dict[NodeType, str]] = None,
187
+ edge_cols: Optional[Dict[EdgeType, Tuple[str, str]]] = None,
188
+ **kwargs,
189
+ ) -> HeteroData:
190
+ r"""Functional shortcut to convert relational tables into a :class:`k3_node.data.HeteroData` graph."""
191
+ if id_cols is None:
192
+ raise ValueError("Must provide 'id_cols' mapping node types to their primary key columns.")
193
+ if edges and edge_cols is None:
194
+ raise ValueError("Must provide 'edge_cols' mapping edge types to (source_col, target_col).")
195
+
196
+ etl = RelationalToGraph(
197
+ id_cols=id_cols,
198
+ edge_cols=edge_cols or {},
199
+ **kwargs,
200
+ )
201
+ return etl.fit_transform(nodes, edges)