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,244 @@
1
+ """Table-to-Graph ETL converter for single tabular datasets."""
2
+
3
+ from typing import Any, Callable, Dict, List, Optional, Sequence, Union
4
+ import numpy as np
5
+
6
+ from k3_node.data import Data
7
+ from k3_node.etl.encoders import TabularEncoder, _get_column_names, _get_column_values
8
+ from k3_node.etl.graph_builders import (
9
+ KNNGraphBuilder,
10
+ SimilarityGraphBuilder,
11
+ SharedEntityGraphBuilder,
12
+ SequentialGraphBuilder,
13
+ )
14
+
15
+
16
+ class TableToGraph:
17
+ r"""ETL pipeline converting tabular datasets (DataFrames, CSVs, or dictionaries)
18
+ into graph :class:`k3_node.data.Data` objects with node features, graph topology,
19
+ labels, and split masks.
20
+
21
+ Args:
22
+ feature_cols: List of column names used as node features. If :obj:`None`, all
23
+ columns except :obj:`target_col` and :obj:`id_col` are used.
24
+ target_col (str, optional): Target column used to create ground-truth labels :obj:`y`.
25
+ id_col (str, optional): Unique row identifier column (e.g. customer_id, product_id).
26
+ Stored in :obj:`data.node_ids` and :obj:`data.id_to_index`.
27
+ edge_strategy (str or callable): Strategy for constructing edges (``"knn"``,
28
+ ``"similarity"``, ``"shared_entity"``, ``"sequential"``, or a custom callable).
29
+ (default: ``"knn"``)
30
+ edge_kwargs (dict, optional): Keyword arguments forwarded to the edge builder.
31
+ (e.g., ``{"k": 5, "metric": "cosine"}`` for KNN).
32
+ column_encoders (dict, optional): Explicit mapping from column names to encoder instances.
33
+ train_ratio (float, optional): Fraction of nodes for training mask. (default: ``None``)
34
+ val_ratio (float, optional): Fraction of nodes for validation mask. (default: ``None``)
35
+ test_ratio (float, optional): Fraction of nodes for test mask. (default: ``None``)
36
+ random_state (int, optional): Random seed for mask splits. (default: ``42``)
37
+ """
38
+
39
+ def __init__(
40
+ self,
41
+ feature_cols: Optional[List[str]] = None,
42
+ target_col: Optional[str] = None,
43
+ id_col: Optional[str] = None,
44
+ edge_strategy: Union[str, Callable] = "knn",
45
+ edge_kwargs: Optional[Dict[str, Any]] = None,
46
+ column_encoders: Optional[Dict[str, Any]] = None,
47
+ train_ratio: Optional[float] = None,
48
+ val_ratio: Optional[float] = None,
49
+ test_ratio: Optional[float] = None,
50
+ random_state: int = 42,
51
+ ):
52
+ self.feature_cols = feature_cols
53
+ self.target_col = target_col
54
+ self.id_col = id_col
55
+ self.edge_strategy = edge_strategy
56
+ self.edge_kwargs = edge_kwargs or {}
57
+ self.column_encoders = column_encoders or {}
58
+ self.train_ratio = train_ratio
59
+ self.val_ratio = val_ratio
60
+ self.test_ratio = test_ratio
61
+ self.random_state = random_state
62
+
63
+ self.encoder_ = TabularEncoder(column_encoders=self.column_encoders)
64
+ self.edge_builder_ = self._resolve_edge_builder()
65
+
66
+ def _resolve_edge_builder(self) -> Callable:
67
+ if callable(self.edge_strategy):
68
+ return self.edge_strategy
69
+
70
+ strategy = str(self.edge_strategy).lower()
71
+ if strategy == "knn":
72
+ return KNNGraphBuilder(**self.edge_kwargs)
73
+ elif strategy in ("similarity", "sim"):
74
+ return SimilarityGraphBuilder(**self.edge_kwargs)
75
+ elif strategy in ("shared_entity", "shared", "bipartite"):
76
+ return SharedEntityGraphBuilder(**self.edge_kwargs)
77
+ elif strategy in ("sequential", "sequence", "temporal"):
78
+ return SequentialGraphBuilder(**self.edge_kwargs)
79
+ else:
80
+ raise ValueError(
81
+ f"Unknown edge_strategy '{self.edge_strategy}'. "
82
+ f"Available: 'knn', 'similarity', 'shared_entity', 'sequential', or custom callable."
83
+ )
84
+
85
+ def fit(self, df_or_dict: Any):
86
+ r"""Fits the tabular feature encoder on the input table."""
87
+ all_cols = _get_column_names(df_or_dict)
88
+ ignore_cols = set()
89
+ if self.target_col:
90
+ ignore_cols.add(self.target_col)
91
+ if self.id_col:
92
+ ignore_cols.add(self.id_col)
93
+
94
+ if self.feature_cols is None:
95
+ fit_cols = [c for c in all_cols if c not in ignore_cols]
96
+ else:
97
+ fit_cols = [c for c in self.feature_cols if c in all_cols and c not in ignore_cols]
98
+
99
+ self.encoder_.fit(df_or_dict, columns=fit_cols)
100
+ return self
101
+
102
+ def transform(self, df_or_dict: Any) -> Data:
103
+ r"""Encodes features, constructs graph edges, and builds a :class:`Data` object."""
104
+ # 1. Node features
105
+ x = self.encoder_.transform(df_or_dict)
106
+ num_nodes = x.shape[0]
107
+
108
+ # 2. Graph topology
109
+ edge_index, edge_attr = self.edge_builder_(x=x, df_or_dict=df_or_dict)
110
+
111
+ # 3. Target labels y
112
+ y = None
113
+ if self.target_col is not None:
114
+ raw_y = _get_column_values(df_or_dict, self.target_col)
115
+ # Check if categorical or numerical
116
+ if all(isinstance(v, (int, np.integer)) for v in raw_y if v is not None):
117
+ y = np.array(raw_y, dtype=np.int64)
118
+ elif all(isinstance(v, (float, int, np.floating, np.integer)) for v in raw_y if v is not None):
119
+ y = np.array(raw_y, dtype=np.float32)
120
+ else:
121
+ # String labels -> label encode to ints
122
+ unique_classes = sorted(list(set(raw_y)))
123
+ mapping = {c: i for i, c in enumerate(unique_classes)}
124
+ y = np.array([mapping[c] for c in raw_y], dtype=np.int64)
125
+
126
+ data = Data(x=x, edge_index=edge_index, edge_attr=edge_attr, y=y)
127
+
128
+ # 4. Optional entity ID tracking
129
+ if self.id_col is not None:
130
+ raw_ids = _get_column_values(df_or_dict, self.id_col)
131
+ data.node_ids = list(raw_ids)
132
+ data.id_to_index = {raw_id: idx for idx, raw_id in enumerate(raw_ids)}
133
+ data.index_to_id = {idx: raw_id for idx, raw_id in enumerate(raw_ids)}
134
+
135
+ # 5. Split masks
136
+ if self.train_ratio is not None:
137
+ np.random.seed(self.random_state)
138
+ indices = np.random.permutation(num_nodes)
139
+
140
+ n_train = int(num_nodes * self.train_ratio)
141
+ val_ratio = self.val_ratio or 0.0
142
+ n_val = int(num_nodes * val_ratio)
143
+
144
+ train_idx = indices[:n_train]
145
+ val_idx = indices[n_train : n_train + n_val]
146
+ test_idx = indices[n_train + n_val :]
147
+
148
+ train_mask = np.zeros(num_nodes, dtype=bool)
149
+ train_mask[train_idx] = True
150
+ data.train_mask = train_mask
151
+
152
+ if val_ratio > 0:
153
+ val_mask = np.zeros(num_nodes, dtype=bool)
154
+ val_mask[val_idx] = True
155
+ data.val_mask = val_mask
156
+
157
+ if self.test_ratio is not None or len(test_idx) > 0:
158
+ test_mask = np.zeros(num_nodes, dtype=bool)
159
+ test_mask[test_idx] = True
160
+ data.test_mask = test_mask
161
+
162
+ return data
163
+
164
+ def fit_transform(self, df_or_dict: Any) -> Data:
165
+ r"""Fits encoders and transforms tabular data into a graph in a single call."""
166
+ return self.fit(df_or_dict).transform(df_or_dict)
167
+
168
+ @classmethod
169
+ def from_dataframe(
170
+ cls,
171
+ df: Any,
172
+ feature_cols: Optional[List[str]] = None,
173
+ target_col: Optional[str] = None,
174
+ id_col: Optional[str] = None,
175
+ edge_strategy: Union[str, Callable] = "knn",
176
+ **kwargs,
177
+ ) -> Data:
178
+ r"""Convenience factory method directly converting a pandas DataFrame into a :class:`Data` object."""
179
+ etl = cls(
180
+ feature_cols=feature_cols,
181
+ target_col=target_col,
182
+ id_col=id_col,
183
+ edge_strategy=edge_strategy,
184
+ **kwargs,
185
+ )
186
+ return etl.fit_transform(df)
187
+
188
+ @classmethod
189
+ def from_csv(
190
+ cls,
191
+ filepath: str,
192
+ feature_cols: Optional[List[str]] = None,
193
+ target_col: Optional[str] = None,
194
+ id_col: Optional[str] = None,
195
+ edge_strategy: Union[str, Callable] = "knn",
196
+ **kwargs,
197
+ ) -> Data:
198
+ r"""Convenience factory method directly loading a CSV file and converting it into a :class:`Data` object."""
199
+ import csv
200
+
201
+ # Parse CSV into dict of column lists
202
+ with open(filepath, mode="r", encoding="utf-8") as f:
203
+ reader = csv.DictReader(f)
204
+ data_dict: Dict[str, List[Any]] = {field: [] for field in reader.fieldnames or []}
205
+ for row in reader:
206
+ for k, v in row.items():
207
+ # Attempt numeric cast
208
+ try:
209
+ v = float(v) if "." in v else int(v)
210
+ except ValueError:
211
+ pass
212
+ data_dict[k].append(v)
213
+
214
+ return cls.from_dataframe(
215
+ df=data_dict,
216
+ feature_cols=feature_cols,
217
+ target_col=target_col,
218
+ id_col=id_col,
219
+ edge_strategy=edge_strategy,
220
+ **kwargs,
221
+ )
222
+
223
+
224
+ # Alias
225
+ TabularToGraph = TableToGraph
226
+
227
+
228
+ def table_to_graph(
229
+ df_or_dict: Any,
230
+ feature_cols: Optional[List[str]] = None,
231
+ target_col: Optional[str] = None,
232
+ id_col: Optional[str] = None,
233
+ edge_strategy: Union[str, Callable] = "knn",
234
+ **kwargs,
235
+ ) -> Data:
236
+ r"""Functional shortcut to convert a table or dictionary into a :class:`k3_node.data.Data` object."""
237
+ etl = TableToGraph(
238
+ feature_cols=feature_cols,
239
+ target_col=target_col,
240
+ id_col=id_col,
241
+ edge_strategy=edge_strategy,
242
+ **kwargs,
243
+ )
244
+ return etl.fit_transform(df_or_dict)
@@ -0,0 +1,318 @@
1
+ """Unit tests for K3-Node Tabular-to-Graph ETL pipelines."""
2
+
3
+ import os
4
+ import tempfile
5
+ import pytest
6
+ import numpy as np
7
+ import pandas as pd
8
+ from keras import ops
9
+
10
+ import k3_node
11
+ from k3_node.data import Data, HeteroData
12
+ from k3_node.etl import (
13
+ NumericalEncoder,
14
+ CategoricalEncoder,
15
+ TabularEncoder,
16
+ KNNGraphBuilder,
17
+ SimilarityGraphBuilder,
18
+ SharedEntityGraphBuilder,
19
+ SequentialGraphBuilder,
20
+ TableToGraph,
21
+ TabularToGraph,
22
+ table_to_graph,
23
+ RelationalToGraph,
24
+ relational_to_graph,
25
+ )
26
+ from k3_node.tasks import NodeClassifier
27
+
28
+
29
+ # ==============================================================================
30
+ # 1. Encoders Unit Tests
31
+ # ==============================================================================
32
+
33
+ def test_numerical_encoder():
34
+ raw = [1.0, 2.0, 3.0, np.nan, 5.0]
35
+
36
+ # Standard scaling
37
+ enc_std = NumericalEncoder(strategy="standard", impute_strategy="mean")
38
+ res_std = enc_std.fit_transform(raw)
39
+ assert res_std.shape == (5, 1)
40
+ assert not np.isnan(res_std).any()
41
+
42
+ # MinMax scaling
43
+ enc_mm = NumericalEncoder(strategy="minmax", impute_strategy="zero")
44
+ res_mm = enc_mm.fit_transform(raw)
45
+ assert res_mm.shape == (5, 1)
46
+ assert res_mm.min() >= 0.0 and res_mm.max() <= 1.0
47
+
48
+ # Log1p
49
+ enc_log = NumericalEncoder(strategy="log1p")
50
+ res_log = enc_log.fit_transform([0.0, 1.0, 10.0])
51
+ assert np.allclose(res_log[0, 0], 0.0)
52
+
53
+
54
+ def test_categorical_encoder():
55
+ raw = ["apple", "banana", "apple", None, "orange"]
56
+
57
+ # One-hot
58
+ enc_oh = CategoricalEncoder(strategy="onehot", handle_unknown="ignore")
59
+ res_oh = enc_oh.fit_transform(raw)
60
+ assert res_oh.shape[0] == 5
61
+ assert res_oh.shape[1] >= 3
62
+
63
+ # Unknown category handling
64
+ transformed = enc_oh.transform(["grape", "banana"])
65
+ assert transformed.shape == (2, res_oh.shape[1])
66
+ assert transformed[0].sum() == 0.0 # unknown ignored
67
+ assert transformed[1].sum() == 1.0 # banana matched
68
+
69
+ # Ordinal
70
+ enc_ord = CategoricalEncoder(strategy="ordinal", unknown_value=-1)
71
+ res_ord = enc_ord.fit_transform(raw)
72
+ assert res_ord.shape == (5, 1)
73
+ assert res_ord.dtype == np.int64
74
+
75
+ # Hash
76
+ enc_hash = CategoricalEncoder(strategy="hash", hash_dim=8)
77
+ res_hash = enc_hash.fit_transform(raw)
78
+ assert res_hash.shape == (5, 8)
79
+
80
+
81
+ def test_tabular_encoder():
82
+ df = pd.DataFrame({
83
+ "age": [25, 30, 35, 40],
84
+ "salary": [50000.0, 60000.0, 75000.0, 90000.0],
85
+ "city": ["NY", "SF", "NY", "LA"],
86
+ })
87
+
88
+ encoder = TabularEncoder()
89
+ x = encoder.fit_transform(df)
90
+ assert x.shape[0] == 4
91
+ # age (1) + salary (1) + city (3 one-hot) = 5
92
+ assert x.shape[1] == 5
93
+ assert x.dtype == np.float32
94
+
95
+
96
+ # ==============================================================================
97
+ # 2. Graph Builders Unit Tests
98
+ # ==============================================================================
99
+
100
+ def test_knn_graph_builder():
101
+ x = np.array([
102
+ [0.0, 0.0],
103
+ [0.1, 0.1],
104
+ [10.0, 10.0],
105
+ [10.1, 10.1],
106
+ ], dtype=np.float32)
107
+
108
+ knn = KNNGraphBuilder(k=1, metric="euclidean", loop=False, bidirectional=True)
109
+ edge_index, edge_attr = knn(x)
110
+
111
+ assert edge_index.shape[0] == 2
112
+ assert edge_index.shape[1] > 0
113
+ # Node 0 and Node 1 should be connected
114
+ edges = set(zip(edge_index[0].tolist(), edge_index[1].tolist()))
115
+ assert (0, 1) in edges and (1, 0) in edges
116
+ # Node 2 and Node 3 should be connected
117
+ assert (2, 3) in edges and (3, 2) in edges
118
+
119
+
120
+ def test_similarity_graph_builder():
121
+ x = np.array([
122
+ [1.0, 0.0],
123
+ [0.99, 0.01],
124
+ [0.0, 1.0],
125
+ ], dtype=np.float32)
126
+
127
+ sim = SimilarityGraphBuilder(threshold=0.9, metric="cosine", loop=False)
128
+ edge_index, edge_attr = sim(x)
129
+
130
+ edges = set(zip(edge_index[0].tolist(), edge_index[1].tolist()))
131
+ assert (0, 1) in edges and (1, 0) in edges
132
+ assert (0, 2) not in edges
133
+
134
+
135
+ def test_shared_entity_graph_builder():
136
+ df = {
137
+ "user_id": ["u1", "u2", "u3", "u4"],
138
+ "device_id": ["d1", "d1", "d2", "d2"],
139
+ }
140
+ x = np.zeros((4, 2), dtype=np.float32)
141
+
142
+ builder = SharedEntityGraphBuilder(entity_cols=["device_id"], loop=False)
143
+ edge_index, edge_attr = builder(x, df_or_dict=df)
144
+
145
+ edges = set(zip(edge_index[0].tolist(), edge_index[1].tolist()))
146
+ # u1 (0) and u2 (1) share d1
147
+ assert (0, 1) in edges and (1, 0) in edges
148
+ # u3 (2) and u4 (3) share d2
149
+ assert (2, 3) in edges and (3, 2) in edges
150
+ # u1 (0) and u3 (2) do not share
151
+ assert (0, 2) not in edges
152
+
153
+
154
+ def test_sequential_graph_builder():
155
+ df = {
156
+ "timestamp": [100, 102, 101, 200, 201],
157
+ "session": ["s1", "s1", "s1", "s2", "s2"],
158
+ }
159
+ x = np.zeros((5, 2), dtype=np.float32)
160
+
161
+ seq = SequentialGraphBuilder(order_col="timestamp", group_by_col="session", window_size=1)
162
+ edge_index, edge_attr = seq(x, df_or_dict=df)
163
+
164
+ edges = set(zip(edge_index[0].tolist(), edge_index[1].tolist()))
165
+ # s1 sequence ordered: 0 (ts=100) -> 2 (ts=101) -> 1 (ts=102)
166
+ assert (0, 2) in edges
167
+ assert (2, 1) in edges
168
+ # s2 sequence: 3 (ts=200) -> 4 (ts=201)
169
+ assert (3, 4) in edges
170
+ # Cross session edges must not exist
171
+ assert (1, 3) not in edges
172
+
173
+
174
+ # ==============================================================================
175
+ # 3. TableToGraph (Single-Table ETL) Tests
176
+ # ==============================================================================
177
+
178
+ def test_table_to_graph_dataframe():
179
+ df = pd.DataFrame({
180
+ "customer_id": ["c1", "c2", "c3", "c4", "c5", "c6"],
181
+ "age": [20, 25, 30, 45, 50, 55],
182
+ "spend": [100.0, 150.0, 200.0, 500.0, 550.0, 600.0],
183
+ "category": ["retail", "retail", "tech", "retail", "tech", "tech"],
184
+ "churn": [0, 0, 0, 1, 1, 1],
185
+ })
186
+
187
+ etl = TableToGraph(
188
+ target_col="churn",
189
+ id_col="customer_id",
190
+ edge_strategy="knn",
191
+ edge_kwargs={"k": 2},
192
+ train_ratio=0.5,
193
+ test_ratio=0.5,
194
+ )
195
+ data = etl.fit_transform(df)
196
+
197
+ assert isinstance(data, Data)
198
+ assert data.x.shape[0] == 6
199
+ assert data.edge_index.shape[0] == 2
200
+ assert data.edge_index.shape[1] > 0
201
+ assert data.y.shape[0] == 6
202
+ assert data.train_mask.shape[0] == 6
203
+ assert data.test_mask.shape[0] == 6
204
+ assert len(data.node_ids) == 6
205
+ assert data.id_to_index["c1"] == 0
206
+
207
+
208
+ def test_table_to_graph_functional():
209
+ data_dict = {
210
+ "feat1": [1.0, 2.0, 3.0, 4.0],
211
+ "feat2": [4.0, 3.0, 2.0, 1.0],
212
+ "label": ["A", "B", "A", "B"],
213
+ }
214
+ data = table_to_graph(data_dict, target_col="label", edge_strategy="knn", edge_kwargs={"k": 1})
215
+ assert isinstance(data, Data)
216
+ assert data.x.shape == (4, 2)
217
+ assert data.y.dtype == np.int64
218
+
219
+
220
+ def test_table_to_graph_from_csv():
221
+ with tempfile.TemporaryDirectory() as tmpdir:
222
+ csv_path = os.path.join(tmpdir, "sample.csv")
223
+ df = pd.DataFrame({
224
+ "id": ["n1", "n2", "n3", "n4"],
225
+ "v1": [10.0, 20.0, 30.0, 40.0],
226
+ "v2": [1.0, 2.0, 3.0, 4.0],
227
+ "y": [0, 1, 0, 1],
228
+ })
229
+ df.to_csv(csv_path, index=False)
230
+
231
+ data = TableToGraph.from_csv(csv_path, target_col="y", id_col="id", edge_strategy="knn", edge_kwargs={"k": 1})
232
+ assert isinstance(data, Data)
233
+ assert data.x.shape[0] == 4
234
+ assert data.node_ids == ["n1", "n2", "n3", "n4"]
235
+
236
+
237
+ # ==============================================================================
238
+ # 4. RelationalToGraph (Multi-Table ETL) Tests
239
+ # ==============================================================================
240
+
241
+ def test_relational_to_graph():
242
+ users_df = pd.DataFrame({
243
+ "user_id": ["u101", "u102", "u103"],
244
+ "age": [22, 35, 48],
245
+ "segment": ["bronze", "gold", "silver"],
246
+ })
247
+ products_df = pd.DataFrame({
248
+ "prod_id": ["p1", "p2", "p3", "p4"],
249
+ "price": [9.99, 49.99, 19.99, 99.99],
250
+ "department": ["books", "elec", "books", "elec"],
251
+ })
252
+ purchases_df = pd.DataFrame({
253
+ "u_id": ["u101", "u101", "u102", "u103", "unknown_user"],
254
+ "p_id": ["p1", "p2", "p3", "p4", "p1"],
255
+ "rating": [5.0, 4.0, 5.0, 3.0, 1.0],
256
+ })
257
+
258
+ etl = RelationalToGraph(
259
+ id_cols={"user": "user_id", "product": "prod_id"},
260
+ edge_cols={("user", "buys", "product"): ("u_id", "p_id")},
261
+ edge_attr_cols={("user", "buys", "product"): ["rating"]},
262
+ )
263
+
264
+ hetero_data = etl.fit_transform(
265
+ nodes={"user": users_df, "product": products_df},
266
+ edges={("user", "buys", "product"): purchases_df},
267
+ )
268
+
269
+ assert isinstance(hetero_data, HeteroData)
270
+ # Check node tables
271
+ assert hetero_data["user"].x.shape[0] == 3
272
+ assert hetero_data["product"].x.shape[0] == 4
273
+
274
+ # Check edges (unknown_user should be filtered out)
275
+ edge_index = hetero_data["user", "buys", "product"].edge_index
276
+ assert edge_index.shape[0] == 2
277
+ assert edge_index.shape[1] == 4 # 4 valid purchases
278
+
279
+ # Check edge attributes
280
+ edge_attr = hetero_data["user", "buys", "product"].edge_attr
281
+ assert edge_attr.shape == (4, 1)
282
+
283
+ # Check ID mapping consistency
284
+ assert hetero_data.id_maps["user"]["u101"] == 0
285
+ assert hetero_data.inverse_id_maps["user"][0] == "u101"
286
+
287
+
288
+ # ==============================================================================
289
+ # 5. End-to-End ETL to GNN Training Integration
290
+ # ==============================================================================
291
+
292
+ def test_etl_to_node_classifier_integration():
293
+ """Verify that a DataFrame converted via TableToGraph can directly train a NodeClassifier."""
294
+ df = pd.DataFrame({
295
+ "feat_a": np.random.randn(30).astype("float32"),
296
+ "feat_b": np.random.randn(30).astype("float32"),
297
+ "category": (["cat1", "cat2", "cat3"] * 10),
298
+ "target": (np.random.randint(0, 2, size=30)),
299
+ })
300
+
301
+ # 1. ETL: Tabular -> Graph Data
302
+ data = table_to_graph(
303
+ df,
304
+ target_col="target",
305
+ edge_strategy="knn",
306
+ edge_kwargs={"k": 3},
307
+ train_ratio=0.7,
308
+ test_ratio=0.3,
309
+ )
310
+
311
+ # 2. Train NodeClassifier in 3 lines of code!
312
+ clf = NodeClassifier(backbone="gcn", hidden_channels=16, num_layers=2)
313
+ clf.fit(data, epochs=2, verbose=0)
314
+
315
+ # 3. Predict & evaluate
316
+ metrics = clf.evaluate(data, mask="test_mask")
317
+ assert "accuracy" in metrics
318
+ assert 0.0 <= metrics["accuracy"] <= 1.0
@@ -0,0 +1,15 @@
1
+ """K3-Node turn-key export and serving module for ONNX, TensorRT, and TensorFlow Lite."""
2
+
3
+ from k3_node.export.onnx_exporter import export_onnx
4
+ from k3_node.export.tflite_exporter import export_tflite
5
+ from k3_node.export.tensorrt_exporter import export_tensorrt, generate_triton_config
6
+ from k3_node.export.runtime import ONNXModel, TFLiteModel
7
+
8
+ __all__ = [
9
+ "export_onnx",
10
+ "export_tflite",
11
+ "export_tensorrt",
12
+ "generate_triton_config",
13
+ "ONNXModel",
14
+ "TFLiteModel",
15
+ ]