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,172 @@
1
+ """Cross-backend export bridge for PyTorch and JAX Keras backends."""
2
+
3
+ import importlib
4
+ import json
5
+ import os
6
+ from pathlib import Path
7
+ import subprocess
8
+ import sys
9
+ import tempfile
10
+ from typing import Any, Dict, List, Optional, Union
11
+ import numpy as np
12
+
13
+ import keras
14
+ from keras import ops
15
+
16
+
17
+ def is_tensorflow_backend() -> bool:
18
+ r"""Checks whether the currently active Keras backend is TensorFlow."""
19
+ return keras.backend.backend() == "tensorflow"
20
+
21
+
22
+ def export_via_tf_subprocess(
23
+ exporter_type: str, # "onnx" or "tflite"
24
+ model_or_task: Any,
25
+ output_path: Union[str, Path],
26
+ dummy_inputs: Optional[Any] = None,
27
+ **kwargs: Any,
28
+ ) -> Path:
29
+ r"""Bridges model export to a TensorFlow worker subprocess when running under PyTorch or JAX.
30
+
31
+ Args:
32
+ exporter_type: "onnx" or "tflite".
33
+ model_or_task: Model or task instance.
34
+ output_path: Target path for the exported model file.
35
+ dummy_inputs: Optional dummy inputs.
36
+ **kwargs: Extra arguments forwarded to the exporter.
37
+
38
+ Returns:
39
+ Path to the exported model file.
40
+ """
41
+ out_file = Path(output_path).resolve()
42
+ out_file.parent.mkdir(parents=True, exist_ok=True)
43
+
44
+ with tempfile.TemporaryDirectory() as tmpdir:
45
+ tmp_path = Path(tmpdir)
46
+ model_dir = tmp_path / "model"
47
+ model_dir.mkdir(parents=True, exist_ok=True)
48
+
49
+ # 1. Ensure model is built before saving weights
50
+ raw_model = getattr(model_or_task, "model", None) or model_or_task
51
+ if hasattr(raw_model, "built") and not raw_model.built:
52
+ from k3_node.export.onnx_exporter import _extract_model_and_inputs
53
+ try:
54
+ _, extracted_inputs, _ = _extract_model_and_inputs(model_or_task, dummy_inputs)
55
+ raw_model(*extracted_inputs)
56
+ except Exception:
57
+ pass
58
+
59
+ # 2. Save model or task weights & config
60
+ if hasattr(model_or_task, "save_pretrained"):
61
+ model_or_task.save_pretrained(str(model_dir))
62
+ cls_module = model_or_task.__class__.__module__
63
+ cls_name = model_or_task.__class__.__name__
64
+ elif hasattr(model_or_task, "model") and hasattr(model_or_task.model, "save_pretrained"):
65
+ model_or_task.model.save_pretrained(str(model_dir))
66
+ cls_module = model_or_task.model.__class__.__module__
67
+ cls_name = model_or_task.model.__class__.__name__
68
+ else:
69
+ raise TypeError(
70
+ f"Model or task of type {type(model_or_task)} must support `save_pretrained` "
71
+ "for cross-backend export."
72
+ )
73
+
74
+ # 2. Serialize dummy inputs if provided
75
+ dummy_path = tmp_path / "dummy.npz"
76
+ dummy_meta_path = tmp_path / "dummy_meta.json"
77
+ has_dummy = False
78
+
79
+ if dummy_inputs is not None:
80
+ arrays: Dict[str, np.ndarray] = {}
81
+ meta: Dict[str, Any] = {}
82
+
83
+ if hasattr(dummy_inputs, "x") and hasattr(dummy_inputs, "edge_index"):
84
+ meta["type"] = "data"
85
+ arrays["x"] = np.asarray(ops.convert_to_numpy(dummy_inputs.x))
86
+ arrays["edge_index"] = np.asarray(ops.convert_to_numpy(dummy_inputs.edge_index))
87
+ if hasattr(dummy_inputs, "batch") and dummy_inputs.batch is not None:
88
+ arrays["batch"] = np.asarray(ops.convert_to_numpy(dummy_inputs.batch))
89
+ elif hasattr(dummy_inputs, "z") and hasattr(dummy_inputs, "pos"):
90
+ meta["type"] = "data"
91
+ arrays["z"] = np.asarray(ops.convert_to_numpy(dummy_inputs.z))
92
+ arrays["pos"] = np.asarray(ops.convert_to_numpy(dummy_inputs.pos))
93
+ if hasattr(dummy_inputs, "batch") and dummy_inputs.batch is not None:
94
+ arrays["batch"] = np.asarray(ops.convert_to_numpy(dummy_inputs.batch))
95
+ elif isinstance(dummy_inputs, (tuple, list)):
96
+ meta["type"] = "tuple"
97
+ for i, elem in enumerate(dummy_inputs):
98
+ arrays[f"arr_{i}"] = np.asarray(ops.convert_to_numpy(elem))
99
+ elif isinstance(dummy_inputs, dict):
100
+ meta["type"] = "dict"
101
+ for k, v in dummy_inputs.items():
102
+ if v is not None:
103
+ arrays[str(k)] = np.asarray(ops.convert_to_numpy(v))
104
+ else:
105
+ meta["type"] = "single"
106
+ arrays["arr_0"] = np.asarray(ops.convert_to_numpy(dummy_inputs))
107
+
108
+ np.savez(str(dummy_path), **arrays)
109
+ with open(dummy_meta_path, "w", encoding="utf-8") as f:
110
+ json.dump(meta, f)
111
+ has_dummy = True
112
+
113
+ # 3. Build worker command
114
+ worker_code = f"""
115
+ import os
116
+ os.environ["KERAS_BACKEND"] = "tensorflow"
117
+ import importlib
118
+ import json
119
+ from pathlib import Path
120
+ import numpy as np
121
+
122
+ import k3_node
123
+ from k3_node.data import Data
124
+ from k3_node.export.onnx_exporter import export_onnx
125
+ from k3_node.export.tflite_exporter import export_tflite
126
+
127
+ # Load model / task
128
+ mod = importlib.import_module("{cls_module}")
129
+ cls = getattr(mod, "{cls_name}")
130
+ model_or_task = cls.from_pretrained(r"{model_dir}")
131
+
132
+ # Load dummy inputs
133
+ dummy = None
134
+ has_dummy = {has_dummy}
135
+ if has_dummy:
136
+ with open(r"{dummy_meta_path}", "r") as f:
137
+ meta = json.load(f)
138
+ npz = np.load(r"{dummy_path}")
139
+ t = meta["type"]
140
+ if t == "data":
141
+ kwargs = {{k: npz[k] for k in npz.files}}
142
+ dummy = Data(**kwargs)
143
+ elif t == "tuple":
144
+ dummy = tuple(npz[f"arr_{{i}}"] for i in range(len(npz.files)))
145
+ elif t == "dict":
146
+ dummy = {{k: npz[k] for k in npz.files}}
147
+ else:
148
+ dummy = npz["arr_0"]
149
+
150
+ extra_kwargs = json.loads(r'''{json.dumps(kwargs)}''')
151
+ if "{exporter_type}" == "onnx":
152
+ export_onnx(model_or_task, r"{out_file}", dummy_inputs=dummy, **extra_kwargs)
153
+ elif "{exporter_type}" == "tflite":
154
+ export_tflite(model_or_task, r"{out_file}", dummy_inputs=dummy, **extra_kwargs)
155
+ """
156
+
157
+ env = dict(os.environ, KERAS_BACKEND="tensorflow")
158
+ res = subprocess.run(
159
+ [sys.executable, "-c", worker_code],
160
+ capture_output=True,
161
+ text=True,
162
+ env=env,
163
+ )
164
+
165
+ if res.returncode != 0:
166
+ raise RuntimeError(
167
+ f"Cross-backend export to {exporter_type.upper()} failed with exit code {res.returncode}:\n"
168
+ f"STDOUT: {res.stdout}\n"
169
+ f"STDERR: {res.stderr}"
170
+ )
171
+
172
+ return out_file
@@ -0,0 +1,190 @@
1
+ """Turn-key ONNX exporter for K3-Node models and tasks."""
2
+
3
+ import os
4
+ from pathlib import Path
5
+ from typing import Any, Dict, List, Optional, Tuple, Union
6
+ import numpy as np
7
+
8
+ import keras
9
+ from keras import ops
10
+
11
+
12
+ def _extract_model_and_inputs(
13
+ model_or_task: Any,
14
+ dummy_inputs: Optional[Any] = None,
15
+ ) -> Tuple[Any, Tuple[Any, ...], List[Any]]:
16
+ r"""Extracts the underlying neural network model and resolves dummy inputs / signatures."""
17
+ # 1. Resolve model
18
+ if hasattr(model_or_task, "model") and model_or_task.model is not None:
19
+ raw_model = model_or_task.model
20
+ else:
21
+ raw_model = model_or_task
22
+
23
+ # 2. Extract inputs from dummy_inputs if provided
24
+ if dummy_inputs is not None:
25
+ if hasattr(dummy_inputs, "x") and hasattr(dummy_inputs, "edge_index"):
26
+ inputs = [dummy_inputs.x, dummy_inputs.edge_index]
27
+ names = ["x", "edge_index"]
28
+ if hasattr(dummy_inputs, "batch") and dummy_inputs.batch is not None:
29
+ inputs.append(dummy_inputs.batch)
30
+ names.append("batch")
31
+ return raw_model, tuple(inputs), names
32
+ elif hasattr(dummy_inputs, "z") and hasattr(dummy_inputs, "pos"):
33
+ inputs = [dummy_inputs.z, dummy_inputs.pos]
34
+ names = ["z", "pos"]
35
+ if hasattr(dummy_inputs, "batch") and dummy_inputs.batch is not None:
36
+ inputs.append(dummy_inputs.batch)
37
+ names.append("batch")
38
+ return raw_model, tuple(inputs), names
39
+ elif isinstance(dummy_inputs, (tuple, list)):
40
+ return raw_model, tuple(dummy_inputs), None
41
+ elif isinstance(dummy_inputs, dict):
42
+ return raw_model, (dummy_inputs,), None
43
+ else:
44
+ return raw_model, (dummy_inputs,), None
45
+
46
+ # 3. Auto-infer dummy inputs from model hyperparameters
47
+ cls_name = model_or_task.__class__.__name__
48
+ in_channels = (
49
+ getattr(model_or_task, "in_channels", None)
50
+ or getattr(raw_model, "in_channels", None)
51
+ or 16
52
+ )
53
+
54
+ if cls_name == "NodeClassifier":
55
+ x = ops.zeros((4, in_channels), dtype="float32")
56
+ edge_index = ops.convert_to_tensor([[0, 1, 2, 3], [1, 2, 3, 0]], dtype="int64")
57
+ return raw_model, (x, edge_index), ["x", "edge_index"]
58
+ elif cls_name in ("GraphClassifier", "GraphRegressor"):
59
+ x = ops.zeros((4, in_channels), dtype="float32")
60
+ edge_index = ops.convert_to_tensor([[0, 1, 2, 3], [1, 2, 3, 0]], dtype="int64")
61
+ batch = ops.convert_to_tensor([0, 0, 1, 1], dtype="int64")
62
+ if hasattr(raw_model, "num_graphs"):
63
+ raw_model.num_graphs = 2
64
+ return raw_model, (x, edge_index, batch), ["x", "edge_index", "batch"]
65
+ elif cls_name == "LinkPredictor":
66
+ x = ops.zeros((4, in_channels), dtype="float32")
67
+ edge_index = ops.convert_to_tensor([[0, 1, 2, 3], [1, 2, 3, 0]], dtype="int64")
68
+ label_idx = ops.convert_to_tensor([[0, 1], [1, 2]], dtype="int64")
69
+ return raw_model, ((x, edge_index), label_idx), ["edge_tuple", "edge_label_index"]
70
+ elif cls_name in ("SchNet", "DimeNet", "DimeNetPlusPlus", "ViSNet", "GNNFF"):
71
+ z = ops.convert_to_tensor([1, 6, 8, 1], dtype="int32")
72
+ pos = ops.zeros((4, 3), dtype="float32")
73
+ return raw_model, (z, pos), ["z", "pos"]
74
+ else:
75
+ # Default standard GNN: (x, edge_index)
76
+ x = ops.zeros((4, in_channels), dtype="float32")
77
+ edge_index = ops.convert_to_tensor([[0, 1, 2, 3], [1, 2, 3, 0]], dtype="int64")
78
+ return raw_model, (x, edge_index), ["x", "edge_index"]
79
+
80
+
81
+ def export_onnx(
82
+ model_or_task: Any,
83
+ output_path: Union[str, Path],
84
+ dummy_inputs: Optional[Any] = None,
85
+ opset: int = 17,
86
+ dynamic_axes: bool = True,
87
+ input_names: Optional[List[str]] = None,
88
+ output_names: Optional[List[str]] = None,
89
+ verbose: bool = False,
90
+ ) -> Path:
91
+ r"""Exports a K3-Node GNN model or task to high-performance ONNX format.
92
+
93
+ Supports arbitrary Graph Neural Networks (GCN, GAT, GraphSAGE, GIN, SchNet,
94
+ materials models, and task estimators) with dynamic graph sizing (varying numbers
95
+ of nodes and edges).
96
+
97
+ Args:
98
+ model_or_task: A K3-Node task instance (e.g. `NodeClassifier`, `GraphClassifier`)
99
+ or model instance (e.g. `GCN`, `SchNet`, `CHGNet`).
100
+ output_path: Target path for the `.onnx` file.
101
+ dummy_inputs: Optional sample input data (e.g., PyG `Data` object, tuple of tensors).
102
+ If `None`, automatically generated based on model topology.
103
+ opset: ONNX operator set version. (default: `17`)
104
+ dynamic_axes: Whether node and edge dimensions should be dynamic. (default: `True`)
105
+ input_names: Optional custom names for input tensors.
106
+ output_names: Optional custom names for output tensors.
107
+ verbose: Whether to print verbose export progress. (default: `False`)
108
+
109
+ Returns:
110
+ Path object pointing to the generated `.onnx` file.
111
+ """
112
+ out_file = Path(output_path)
113
+ out_file.parent.mkdir(parents=True, exist_ok=True)
114
+
115
+ from k3_node.export.cross_backend import is_tensorflow_backend, export_via_tf_subprocess
116
+
117
+ if not is_tensorflow_backend():
118
+ return export_via_tf_subprocess(
119
+ exporter_type="onnx",
120
+ model_or_task=model_or_task,
121
+ output_path=output_path,
122
+ dummy_inputs=dummy_inputs,
123
+ opset=opset,
124
+ dynamic_axes=dynamic_axes,
125
+ input_names=input_names,
126
+ output_names=output_names,
127
+ verbose=verbose,
128
+ )
129
+
130
+ try:
131
+ import tf2onnx
132
+ import tensorflow as tf
133
+ except ImportError:
134
+ raise ImportError(
135
+ "The `tf2onnx` and `tensorflow` packages are required for ONNX export. "
136
+ "Install them via `pip install tf2onnx tensorflow onnx`."
137
+ )
138
+
139
+ model, inputs, inferred_names = _extract_model_and_inputs(model_or_task, dummy_inputs)
140
+ names = input_names or inferred_names or [f"input_{i}" for i in range(len(inputs))]
141
+
142
+ # Build input signature with dynamic axes if requested
143
+ signature = []
144
+ for inp, name in zip(inputs, names):
145
+ inp_np = ops.convert_to_numpy(inp)
146
+ dtype = tf.as_dtype(inp_np.dtype)
147
+ if dynamic_axes:
148
+ if inp_np.ndim == 2 and inp_np.shape[0] == 2 and (
149
+ np.issubdtype(inp_np.dtype, np.integer) or "edge" in name.lower()
150
+ ):
151
+ # Edge index: (2, num_edges) -> dynamic num_edges
152
+ shape = (2, None)
153
+ elif inp_np.ndim == 2:
154
+ # Node feature: (num_nodes, in_channels) -> dynamic num_nodes
155
+ shape = (None, inp_np.shape[1])
156
+ elif inp_np.ndim == 1:
157
+ # Vector (batch or z): (num_nodes,) -> dynamic num_nodes
158
+ shape = (None,)
159
+ else:
160
+ shape = tuple(None if i == 0 else s for i, s in enumerate(inp_np.shape))
161
+ else:
162
+ shape = inp_np.shape
163
+ signature.append(tf.TensorSpec(shape=shape, dtype=dtype, name=name))
164
+
165
+ # Define trace function
166
+ @tf.function(input_signature=signature)
167
+ def forward_fn(*tensors):
168
+ return model(*tensors)
169
+
170
+ if verbose:
171
+ print(f"Exporting model to ONNX with input signature: {signature}")
172
+
173
+ # Convert using tf2onnx
174
+ model_proto, _ = tf2onnx.convert.from_function(
175
+ forward_fn,
176
+ input_signature=signature,
177
+ output_path=str(out_file),
178
+ opset=opset,
179
+ )
180
+
181
+ # Validate ONNX graph
182
+ try:
183
+ import onnx
184
+ onnx_model = onnx.load(str(out_file))
185
+ onnx.checker.check_model(onnx_model)
186
+ except Exception as e:
187
+ if verbose:
188
+ print(f"Warning: ONNX validation check returned: {e}")
189
+
190
+ return out_file
@@ -0,0 +1,254 @@
1
+ """Lightweight inference runtime engines for serving ONNX and TFLite models."""
2
+
3
+ from pathlib import Path
4
+ from typing import Any, Dict, List, Optional, Union
5
+ import numpy as np
6
+
7
+
8
+ class ONNXModel:
9
+ r"""High-performance serving wrapper for exported ONNX GNN models.
10
+
11
+ Requires only `onnxruntime` and `numpy`. Completely decoupled from Keras,
12
+ PyTorch, and TensorFlow for lightweight production microservices.
13
+
14
+ Example:
15
+ ```python
16
+ from k3_node.export import ONNXModel
17
+ model = ONNXModel("cora_gcn.onnx")
18
+ preds = model.predict(graph_data)
19
+ ```
20
+ """
21
+
22
+ def __init__(
23
+ self,
24
+ model_path: Union[str, Path],
25
+ providers: Optional[List[str]] = None,
26
+ session_options: Optional[Any] = None,
27
+ ):
28
+ r"""Initializes the ONNX runtime inference session.
29
+
30
+ Args:
31
+ model_path: Path to the `.onnx` model file.
32
+ providers: Execution providers list (e.g. `['CUDAExecutionProvider', 'CPUExecutionProvider']`).
33
+ If `None`, automatically picks the fastest available provider.
34
+ session_options: Optional custom ONNX Runtime SessionOptions.
35
+ """
36
+ try:
37
+ import onnxruntime as ort
38
+ except ImportError:
39
+ raise ImportError(
40
+ "The `onnxruntime` package is required to load and serve ONNX models. "
41
+ "Install it via `pip install onnxruntime` (or `onnxruntime-gpu` for CUDA/TensorRT)."
42
+ )
43
+
44
+ self.model_path = Path(model_path)
45
+ if not self.model_path.exists():
46
+ raise FileNotFoundError(f"ONNX model file not found at: {self.model_path}")
47
+
48
+ if providers is None:
49
+ available = ort.get_available_providers()
50
+ # Prioritize TensorRT, CUDA, CoreML, DirectML, CPU
51
+ priority = [
52
+ "TensorrtExecutionProvider",
53
+ "CUDAExecutionProvider",
54
+ "CoreMLExecutionProvider",
55
+ "DmlExecutionProvider",
56
+ "CPUExecutionProvider",
57
+ ]
58
+ providers = [p for p in priority if p in available]
59
+
60
+ self.session = ort.InferenceSession(
61
+ str(self.model_path),
62
+ sess_options=session_options,
63
+ providers=providers,
64
+ )
65
+ self.input_names = [inp.name for inp in self.session.get_inputs()]
66
+ self.output_names = [out.name for out in self.session.get_outputs()]
67
+
68
+ def predict(self, data: Any = None, *args: Any, **kwargs: Any) -> np.ndarray:
69
+ r"""Runs low-latency inference on graph data.
70
+
71
+ Accepts:
72
+ - PyG / K3-Node `Data` object
73
+ - Dictionary of tensors
74
+ - Positional numpy arrays
75
+
76
+ Args:
77
+ data: Input graph Data, dictionary, or array.
78
+ *args: Additional positional inputs.
79
+
80
+ Returns:
81
+ Numpy array containing model predictions or logits.
82
+ """
83
+ feed_dict = {}
84
+
85
+ # 1. PyG/K3 Data object
86
+ if hasattr(data, "x") and hasattr(data, "edge_index"):
87
+ feed_dict = self._match_inputs({
88
+ "x": np.asarray(data.x, dtype=np.float32),
89
+ "edge_index": np.asarray(data.edge_index, dtype=np.int64),
90
+ "batch": np.asarray(getattr(data, "batch", None), dtype=np.int64) if getattr(data, "batch", None) is not None else None,
91
+ })
92
+ elif hasattr(data, "z") and hasattr(data, "pos"):
93
+ feed_dict = self._match_inputs({
94
+ "z": np.asarray(data.z, dtype=np.int32),
95
+ "pos": np.asarray(data.pos, dtype=np.float32),
96
+ "batch": np.asarray(getattr(data, "batch", None), dtype=np.int32) if getattr(data, "batch", None) is not None else None,
97
+ })
98
+ elif isinstance(data, dict):
99
+ feed_dict = self._match_inputs(data)
100
+ elif isinstance(data, (tuple, list)):
101
+ for name, val in zip(self.input_names, data):
102
+ feed_dict[name] = np.asarray(val)
103
+ elif data is not None:
104
+ all_args = [data, *args]
105
+ for name, val in zip(self.input_names, all_args):
106
+ feed_dict[name] = np.asarray(val)
107
+
108
+ outputs = self.session.run(self.output_names, feed_dict)
109
+ return outputs[0] if len(outputs) == 1 else tuple(outputs)
110
+
111
+ def _match_inputs(self, named_inputs: Dict[str, Any]) -> Dict[str, np.ndarray]:
112
+ matched = {}
113
+ valid_inputs = {k: np.asarray(v) for k, v in named_inputs.items() if v is not None}
114
+ lower_inputs = {k.lower(): v for k, v in valid_inputs.items()}
115
+
116
+ used_keys = set()
117
+ session_inputs = self.session.get_inputs()
118
+
119
+ # Step 1: Match exact or known semantic names
120
+ for sess_inp in session_inputs:
121
+ name = sess_inp.name
122
+ clean = name.split(":")[0].lower()
123
+
124
+ target_val = None
125
+ matched_key = None
126
+
127
+ if clean in lower_inputs:
128
+ target_val = lower_inputs[clean]
129
+ matched_key = clean
130
+ elif "edge" in clean and "edge_index" in lower_inputs:
131
+ target_val = lower_inputs["edge_index"]
132
+ matched_key = "edge_index"
133
+ elif clean == "x" and "x" in lower_inputs:
134
+ target_val = lower_inputs["x"]
135
+ matched_key = "x"
136
+ elif clean in ("batch", "batch_idx") and "batch" in lower_inputs:
137
+ target_val = lower_inputs["batch"]
138
+ matched_key = "batch"
139
+ elif clean in ("z", "atomic_numbers") and "z" in lower_inputs:
140
+ target_val = lower_inputs["z"]
141
+ matched_key = "z"
142
+ elif clean in ("pos", "positions", "coord") and "pos" in lower_inputs:
143
+ target_val = lower_inputs["pos"]
144
+ matched_key = "pos"
145
+
146
+ if target_val is not None:
147
+ matched[name] = target_val
148
+ used_keys.add(matched_key)
149
+
150
+ # Step 2: Positional fallback for remaining inputs
151
+ unmatched_session = [inp for inp in session_inputs if inp.name not in matched]
152
+ unmatched_keys = [k for k in valid_inputs.keys() if k.lower() not in used_keys]
153
+
154
+ if unmatched_session and unmatched_keys:
155
+ for sess_inp, k in zip(unmatched_session, unmatched_keys):
156
+ matched[sess_inp.name] = valid_inputs[k]
157
+
158
+ # Step 3: Align data types with expected ONNX session tensor types
159
+ type_map = {
160
+ "tensor(float)": np.float32,
161
+ "tensor(float16)": np.float16,
162
+ "tensor(double)": np.float64,
163
+ "tensor(int64)": np.int64,
164
+ "tensor(int32)": np.int32,
165
+ "tensor(int8)": np.int8,
166
+ "tensor(uint8)": np.uint8,
167
+ "tensor(bool)": np.bool_,
168
+ }
169
+ for sess_inp in session_inputs:
170
+ if sess_inp.name in matched:
171
+ expected_np_type = type_map.get(sess_inp.type)
172
+ if expected_np_type and matched[sess_inp.name].dtype != expected_np_type:
173
+ matched[sess_inp.name] = matched[sess_inp.name].astype(expected_np_type)
174
+
175
+ return matched
176
+
177
+
178
+ class TFLiteModel:
179
+ r"""Lightweight serving wrapper for TensorFlow Lite flatbuffer GNN models.
180
+
181
+ Requires only standard `tensorflow` or `tflite_runtime`. Ideal for mobile,
182
+ Raspberry Pi, and edge embedded devices.
183
+
184
+ Example:
185
+ ```python
186
+ from k3_node.export import TFLiteModel
187
+ model = TFLiteModel("model.tflite")
188
+ preds = model.predict(graph_data)
189
+ ```
190
+ """
191
+
192
+ def __init__(self, model_path: Union[str, Path]):
193
+ r"""Initializes the TFLite interpreter.
194
+
195
+ Args:
196
+ model_path: Path to the `.tflite` model file.
197
+ """
198
+ self.model_path = Path(model_path)
199
+ if not self.model_path.exists():
200
+ raise FileNotFoundError(f"TFLite model file not found at: {self.model_path}")
201
+
202
+ try:
203
+ import tensorflow as tf
204
+ self.interpreter = tf.lite.Interpreter(model_path=str(self.model_path))
205
+ except ImportError:
206
+ try:
207
+ import tflite_runtime.interpreter as tflite
208
+ self.interpreter = tflite.Interpreter(model_path=str(self.model_path))
209
+ except ImportError:
210
+ raise ImportError(
211
+ "Either `tensorflow` or `tflite_runtime` is required to run TFLite models. "
212
+ "Install via `pip install tflite-runtime` or `pip install tensorflow`."
213
+ )
214
+
215
+ self.interpreter.allocate_tensors()
216
+ self.input_details = self.interpreter.get_input_details()
217
+ self.output_details = self.interpreter.get_output_details()
218
+
219
+ def predict(self, data: Any = None, *args: Any, **kwargs: Any) -> np.ndarray:
220
+ r"""Runs inference using the TFLite interpreter.
221
+
222
+ Args:
223
+ data: Input graph Data, dict, or numpy array.
224
+ *args: Additional positional inputs.
225
+
226
+ Returns:
227
+ Numpy array containing prediction results.
228
+ """
229
+ inputs = []
230
+ if hasattr(data, "x") and hasattr(data, "edge_index"):
231
+ inputs = [np.asarray(data.x, dtype=np.float32), np.asarray(data.edge_index, dtype=np.int64)]
232
+ if hasattr(data, "batch") and data.batch is not None:
233
+ inputs.append(np.asarray(data.batch, dtype=np.int64))
234
+ elif hasattr(data, "z") and hasattr(data, "pos"):
235
+ inputs = [np.asarray(data.z, dtype=np.int32), np.asarray(data.pos, dtype=np.float32)]
236
+ if hasattr(data, "batch") and data.batch is not None:
237
+ inputs.append(np.asarray(data.batch, dtype=np.int32))
238
+ elif isinstance(data, (tuple, list)):
239
+ inputs = [np.asarray(x) for x in data]
240
+ elif isinstance(data, dict):
241
+ inputs = [np.asarray(v) for v in data.values()]
242
+ elif data is not None:
243
+ inputs = [np.asarray(data)] + [np.asarray(a) for a in args]
244
+
245
+ # Feed tensors
246
+ for detail, inp in zip(self.input_details, inputs):
247
+ # Cast dtype to match interpreter expectation
248
+ target_dtype = detail["dtype"]
249
+ inp_cast = inp.astype(target_dtype) if inp.dtype != target_dtype else inp
250
+ self.interpreter.set_tensor(detail["index"], inp_cast)
251
+
252
+ self.interpreter.invoke()
253
+ outputs = [self.interpreter.get_tensor(d["index"]) for d in self.output_details]
254
+ return outputs[0] if len(outputs) == 1 else tuple(outputs)