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,162 @@
1
+ """Verbalization and LLM prompt formatting for GraphRAG."""
2
+
3
+ from typing import Dict, List, Literal, Optional, Sequence, Tuple, Union
4
+ import numpy as np
5
+ from keras import ops
6
+
7
+ from k3_node.rag.subgraph import SubgraphResult
8
+
9
+
10
+ def subgraph_to_triples(
11
+ subgraph: SubgraphResult,
12
+ id_to_entity: Optional[Dict[int, str]] = None,
13
+ id_to_relation: Optional[Dict[int, str]] = None,
14
+ ) -> List[Tuple[str, str, str]]:
15
+ """Convert a SubgraphResult into a list of (head, relation, tail) string triples.
16
+
17
+ Args:
18
+ subgraph: The extracted SubgraphResult.
19
+ id_to_entity: Optional map from original entity integer ID to entity string.
20
+ id_to_relation: Optional map from relation integer ID to relation string.
21
+
22
+ Returns:
23
+ List of `(head, relation, tail)` string triples.
24
+ """
25
+ if subgraph.num_edges == 0:
26
+ return []
27
+
28
+ edge_index_np = np.asarray(ops.convert_to_numpy(subgraph.edge_index)).astype(np.int64)
29
+ row, col = edge_index_np[0], edge_index_np[1]
30
+
31
+ # Map subgraph indices back to original node IDs if mapping/nodes available
32
+ if subgraph.nodes is not None and len(subgraph.nodes) > 0:
33
+ orig_row = subgraph.nodes[row]
34
+ orig_col = subgraph.nodes[col]
35
+ else:
36
+ orig_row, orig_col = row, col
37
+
38
+ edge_type_np = None
39
+ if subgraph.edge_type is not None:
40
+ edge_type_np = np.asarray(ops.convert_to_numpy(subgraph.edge_type)).astype(np.int64)
41
+
42
+ triples = []
43
+ for i in range(len(row)):
44
+ h_id = int(orig_row[i])
45
+ t_id = int(orig_col[i])
46
+ h_name = id_to_entity.get(h_id, f"Node_{h_id}") if id_to_entity else f"Node_{h_id}"
47
+ t_name = id_to_entity.get(t_id, f"Node_{t_id}") if id_to_entity else f"Node_{t_id}"
48
+
49
+ if edge_type_np is not None:
50
+ r_id = int(edge_type_np[i])
51
+ r_name = id_to_relation.get(r_id, f"rel_{r_id}") if id_to_relation else f"rel_{r_id}"
52
+ else:
53
+ r_name = "connected_to"
54
+
55
+ triples.append((h_name, r_name, t_name))
56
+
57
+ return triples
58
+
59
+
60
+ def verbalize_subgraph(
61
+ subgraph: SubgraphResult,
62
+ id_to_entity: Optional[Dict[int, str]] = None,
63
+ id_to_relation: Optional[Dict[int, str]] = None,
64
+ format_style: Literal["triples", "markdown", "natural"] = "markdown",
65
+ max_triples: Optional[int] = 50,
66
+ ) -> str:
67
+ """Verbalize an extracted subgraph into textual knowledge for LLM prompt augmentation.
68
+
69
+ Args:
70
+ subgraph: Extracted SubgraphResult around retrieved entities.
71
+ id_to_entity: Optional dictionary mapping node ID to entity name.
72
+ id_to_relation: Optional dictionary mapping relation ID to relation name.
73
+ format_style:
74
+ - `"triples"`: List of `(Head, Relation, Tail)` text triples.
75
+ - `"markdown"`: Markdown bullet list with facts.
76
+ - `"natural"`: Natural language sentences (`"Head relation Tail."`).
77
+ max_triples: Maximum number of facts to include in the context.
78
+
79
+ Returns:
80
+ Formatted textual context string ready to be injected into an LLM prompt.
81
+ """
82
+ triples = subgraph_to_triples(subgraph, id_to_entity, id_to_relation)
83
+ if not triples:
84
+ return "No relevant knowledge graph facts retrieved."
85
+
86
+ if max_triples is not None and len(triples) > max_triples:
87
+ triples = triples[:max_triples]
88
+
89
+ if format_style == "triples":
90
+ lines = [f"({h}, {r}, {t})" for h, r, t in triples]
91
+ return "\n".join(lines)
92
+
93
+ elif format_style == "natural":
94
+ lines = []
95
+ for h, r, t in triples:
96
+ # Clean relation string (replace underscores with spaces)
97
+ rel_str = r.replace("_", " ")
98
+ lines.append(f"{h} {rel_str} {t}.")
99
+ return " ".join(lines)
100
+
101
+ else: # "markdown"
102
+ lines = ["### Retrieved Knowledge Graph Facts:"]
103
+ for h, r, t in triples:
104
+ lines.append(f"- **{h}** — *{r}* -> **{t}**")
105
+ return "\n".join(lines)
106
+
107
+
108
+ def format_llm_prompt(
109
+ query: str,
110
+ context: str,
111
+ system_prompt: Optional[str] = None,
112
+ model_family: Literal["llama3", "mistral", "chatml", "standard"] = "llama3",
113
+ ) -> str:
114
+ """Format query and retrieved KG context into prompt templates for LLMs.
115
+
116
+ Supported model families include:
117
+ - `"llama3"`: Meta Llama 3 / 3.1 instruct template.
118
+ - `"mistral"`: Mistral / Mixtral instruct template.
119
+ - `"chatml"`: OpenAI / Qwen ChatML template.
120
+ - `"standard"`: General markdown system/user format.
121
+
122
+ Args:
123
+ query: The user's input question or instruction.
124
+ context: The verbalized knowledge graph context.
125
+ system_prompt: Optional system prompt to instruct the LLM.
126
+ model_family: Target model prompt format. (default: "llama3")
127
+
128
+ Returns:
129
+ Formatted prompt string.
130
+ """
131
+ if system_prompt is None:
132
+ system_prompt = (
133
+ "You are an expert assistant augmented with a Knowledge Graph. "
134
+ "Use the provided Knowledge Graph Context to accurately answer the user's question."
135
+ )
136
+
137
+ if model_family == "llama3":
138
+ return (
139
+ "<|begin_of_text|><|start_header_id|>system<|end_header_id|>\n\n"
140
+ f"{system_prompt}<|eot_id|><|start_header_id|>user<|end_header_id|>\n\n"
141
+ f"Knowledge Graph Context:\n{context}\n\n"
142
+ f"Question: {query}<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n\n"
143
+ )
144
+ elif model_family == "mistral":
145
+ return (
146
+ f"<s>[INST] {system_prompt}\n\n"
147
+ f"Knowledge Graph Context:\n{context}\n\n"
148
+ f"Question: {query} [/INST]"
149
+ )
150
+ elif model_family == "chatml":
151
+ return (
152
+ f"<|im_start|>system\n{system_prompt}<|im_end|>\n"
153
+ f"<|im_start|>user\nKnowledge Graph Context:\n{context}\n\nQuestion: {query}<|im_end|>\n"
154
+ f"<|im_start|>assistant\n"
155
+ )
156
+ else: # "standard"
157
+ return (
158
+ f"System: {system_prompt}\n\n"
159
+ f"Knowledge Graph Context:\n{context}\n\n"
160
+ f"User Question: {query}\n\n"
161
+ f"Assistant:"
162
+ )
@@ -0,0 +1,19 @@
1
+ """High-level Task APIs for K3-Node."""
2
+
3
+ from k3_node.tasks.base import BaseTask
4
+ from k3_node.tasks.backbone_resolver import resolve_backbone
5
+ from k3_node.tasks.node_classification import NodeClassifier
6
+ from k3_node.tasks.node_regression import NodeRegressor
7
+ from k3_node.tasks.graph_classification import GraphClassifier
8
+ from k3_node.tasks.graph_regression import GraphRegressor
9
+ from k3_node.tasks.link_prediction import LinkPredictor
10
+
11
+ __all__ = [
12
+ "BaseTask",
13
+ "resolve_backbone",
14
+ "NodeClassifier",
15
+ "NodeRegressor",
16
+ "GraphClassifier",
17
+ "GraphRegressor",
18
+ "LinkPredictor",
19
+ ]
@@ -0,0 +1,125 @@
1
+ """Backbone resolver mapping string identifiers to K3-Node models."""
2
+
3
+ from typing import Any, Dict, Optional, Union
4
+ import keras
5
+
6
+ from k3_node import models
7
+
8
+
9
+ BACKBONE_REGISTRY = {
10
+ "gcn": models.GCN,
11
+ "gat": models.GAT,
12
+ "sage": models.GraphSAGE,
13
+ "graphsage": models.GraphSAGE,
14
+ "gin": models.GIN,
15
+ "pna": models.PNA,
16
+ "edge_cnn": models.EdgeCNN,
17
+ "edgecnn": models.EdgeCNN,
18
+ "mlp": models.MLP,
19
+ "linkx": models.LINKX,
20
+ "pmlp": models.PMLP,
21
+ "sgformer": models.SGFormer,
22
+ "polynormer": models.Polynormer,
23
+ "schnet": models.SchNet,
24
+ "dimenet": models.DimeNet,
25
+ "dimenet++": models.DimeNetPlusPlus,
26
+ "dimenetplusplus": models.DimeNetPlusPlus,
27
+ "attentive_fp": models.AttentiveFP,
28
+ "attentivefp": models.AttentiveFP,
29
+ "graph_unet": models.GraphUNet,
30
+ }
31
+
32
+
33
+ def resolve_backbone(
34
+ backbone: Union[str, keras.Model, Any],
35
+ in_channels: int,
36
+ out_channels: int,
37
+ hidden_channels: int = 64,
38
+ num_layers: int = 2,
39
+ dropout: float = 0.5,
40
+ **kwargs,
41
+ ) -> keras.Model:
42
+ r"""Resolves a backbone string or instance into a compiled/callable Keras model."""
43
+ if isinstance(backbone, keras.Model):
44
+ return backbone
45
+
46
+ if not isinstance(backbone, str):
47
+ raise TypeError(
48
+ f"Expected backbone to be a string name or keras.Model instance, got {type(backbone)}"
49
+ )
50
+
51
+ name = backbone.lower().strip()
52
+ if name not in BACKBONE_REGISTRY:
53
+ available = ", ".join(sorted(BACKBONE_REGISTRY.keys()))
54
+ raise ValueError(f"Unknown backbone '{backbone}'. Available backbones: {available}")
55
+
56
+ cls = BACKBONE_REGISTRY[name]
57
+
58
+ # Handle models with specialized constructors
59
+ if name in ("schnet",):
60
+ return cls(hidden_channels=hidden_channels, num_filters=hidden_channels, num_interactions=num_layers, **kwargs)
61
+ elif name in ("dimenet", "dimenet++", "dimenetplusplus"):
62
+ return cls(hidden_channels=hidden_channels, out_channels=out_channels, num_blocks=num_layers, **kwargs)
63
+ elif name in ("attentive_fp", "attentivefp"):
64
+ edge_dim = kwargs.pop("edge_dim", in_channels)
65
+ num_timesteps = kwargs.pop("num_timesteps", 2)
66
+ return cls(
67
+ in_channels=in_channels,
68
+ hidden_channels=hidden_channels,
69
+ out_channels=out_channels,
70
+ edge_dim=edge_dim,
71
+ num_layers=num_layers,
72
+ num_timesteps=num_timesteps,
73
+ dropout=dropout,
74
+ **kwargs,
75
+ )
76
+ elif name in ("mlp",):
77
+ channel_list = [in_channels] + [hidden_channels] * (num_layers - 1) + [out_channels]
78
+ return cls(channel_list=channel_list, dropout=dropout, **kwargs)
79
+ elif name in ("linkx",):
80
+ num_nodes = kwargs.pop("num_nodes", 1000)
81
+ return cls(
82
+ num_nodes=num_nodes,
83
+ in_channels=in_channels,
84
+ hidden_channels=hidden_channels,
85
+ out_channels=out_channels,
86
+ num_layers=num_layers,
87
+ dropout=dropout,
88
+ **kwargs,
89
+ )
90
+ elif name in ("sgformer",):
91
+ return cls(
92
+ in_channels=in_channels,
93
+ hidden_channels=hidden_channels,
94
+ out_channels=out_channels,
95
+ num_layers=num_layers,
96
+ dropout=dropout,
97
+ **kwargs,
98
+ )
99
+ elif name in ("polynormer",):
100
+ return cls(
101
+ in_channels=in_channels,
102
+ hidden_channels=hidden_channels,
103
+ out_channels=out_channels,
104
+ num_layers=num_layers,
105
+ dropout=dropout,
106
+ **kwargs,
107
+ )
108
+ elif name in ("graph_unet",):
109
+ return cls(
110
+ in_channels=in_channels,
111
+ hidden_channels=hidden_channels,
112
+ out_channels=out_channels,
113
+ depth=num_layers,
114
+ **kwargs,
115
+ )
116
+ else:
117
+ # Standard BasicGNN (GCN, GAT, SAGE, GIN, PNA, EdgeCNN)
118
+ return cls(
119
+ in_channels=in_channels,
120
+ hidden_channels=hidden_channels,
121
+ num_layers=num_layers,
122
+ out_channels=out_channels,
123
+ dropout=dropout,
124
+ **kwargs,
125
+ )
k3_node/tasks/base.py ADDED
@@ -0,0 +1,67 @@
1
+ """Base task abstraction for high-level K3-Node estimators."""
2
+
3
+ from typing import Any, Dict, List, Optional, Tuple, Union
4
+ import keras
5
+ from keras import ops
6
+
7
+ from k3_node.data import BaseData
8
+ from k3_node.hub.hub_mixin import K3NodeHubMixin
9
+
10
+
11
+ class BaseTask(K3NodeHubMixin):
12
+ r"""Abstract base task estimator providing common training, evaluation,
13
+ and serialization workflows.
14
+ """
15
+
16
+ def __init__(self, model: Optional[keras.Model] = None):
17
+ self.model = model
18
+ self._is_compiled = False
19
+
20
+ def compile(
21
+ self,
22
+ optimizer: Optional[Union[str, keras.optimizers.Optimizer]] = None,
23
+ loss: Optional[Any] = None,
24
+ metrics: Optional[List[Any]] = None,
25
+ **kwargs,
26
+ ):
27
+ r"""Configures the task model for training."""
28
+ if self.model is None:
29
+ raise RuntimeError("Model has not been initialized. Call fit() or construct with a model first.")
30
+
31
+ opt = optimizer or keras.optimizers.Adam(learning_rate=0.01)
32
+ self.model.compile(optimizer=opt, loss=loss, metrics=metrics, **kwargs)
33
+ self._is_compiled = True
34
+ return self
35
+
36
+ def summary(self):
37
+ r"""Prints a string summary of the underlying neural network."""
38
+ if self.model is not None:
39
+ return self.model.summary()
40
+ print("Model has not been initialized yet.")
41
+
42
+ def save(self, filepath: str):
43
+ r"""Saves the underlying model weights."""
44
+ if self.model is not None:
45
+ self.model.save(filepath)
46
+ else:
47
+ raise RuntimeError("Cannot save an uninitialized model.")
48
+
49
+ @classmethod
50
+ def load(cls, filepath: str, **kwargs):
51
+ r"""Loads a saved task model from disk."""
52
+ model = keras.models.load_model(filepath, **kwargs)
53
+ instance = cls(model=model)
54
+ instance._is_compiled = True
55
+ return instance
56
+
57
+ def _extract_inputs(self, data: Any):
58
+ r"""Extracts input tensors from a Data object, tuple, or dictionary."""
59
+ if hasattr(data, "inputs"):
60
+ return data.inputs
61
+ elif isinstance(data, (tuple, list)):
62
+ return data
63
+ elif hasattr(data, "x") and hasattr(data, "edge_index"):
64
+ if hasattr(data, "edge_attr") and data.edge_attr is not None:
65
+ return (data.x, data.edge_index, data.edge_attr)
66
+ return (data.x, data.edge_index)
67
+ return data
@@ -0,0 +1,270 @@
1
+ """High-level Graph Classification Task."""
2
+
3
+ from typing import Any, Dict, List, Optional, Union
4
+ import keras
5
+ from keras import layers, ops
6
+ import numpy as np
7
+
8
+ from k3_node.tasks.base import BaseTask
9
+ from k3_node.tasks.backbone_resolver import resolve_backbone
10
+ from k3_node.layers import pool as k3_pool
11
+ from k3_node.loader import DataLoader
12
+
13
+
14
+ class GraphClassificationModel(keras.Model):
15
+ r"""Internal wrapper combining a node-level GNN backbone, a global readout
16
+ pooling operation, and a final classification dense head.
17
+ """
18
+
19
+ def __init__(
20
+ self,
21
+ backbone: keras.Model,
22
+ pooling: str = "mean",
23
+ hidden_channels: int = 64,
24
+ num_classes: int = 2,
25
+ dropout: float = 0.5,
26
+ ):
27
+ super().__init__()
28
+ self.backbone = backbone
29
+ self.pooling = pooling
30
+ self.dropout = layers.Dropout(dropout) if dropout > 0 else None
31
+ self.head = layers.Dense(num_classes)
32
+ self.num_graphs = None
33
+
34
+ def call(self, inputs, training=False):
35
+ if isinstance(inputs, (tuple, list)):
36
+ x, edge_index = inputs[0], inputs[1]
37
+ batch = inputs[2] if len(inputs) > 2 else None
38
+ size = inputs[3] if len(inputs) > 3 else self.num_graphs
39
+ else:
40
+ x = inputs
41
+ edge_index = getattr(x, "edge_index", None)
42
+ batch = getattr(x, "batch", None)
43
+ size = getattr(x, "num_graphs", self.num_graphs)
44
+ x = getattr(x, "x", x)
45
+
46
+ if batch is None:
47
+ batch = ops.zeros((ops.shape(x)[0],), dtype="int64")
48
+
49
+ h = self.backbone((x, edge_index), training=training)
50
+
51
+ if self.pooling in ("mean", "global_mean_pool"):
52
+ g = k3_pool.global_mean_pool(h, batch, size=size)
53
+ elif self.pooling in ("add", "sum", "global_add_pool"):
54
+ g = k3_pool.global_add_pool(h, batch, size=size)
55
+ elif self.pooling in ("max", "global_max_pool"):
56
+ g = k3_pool.global_max_pool(h, batch, size=size)
57
+ else:
58
+ g = k3_pool.global_mean_pool(h, batch, size=size)
59
+
60
+ if self.dropout is not None:
61
+ g = self.dropout(g, training=training)
62
+ return self.head(g)
63
+
64
+
65
+ class GraphClassifier(BaseTask):
66
+ r"""High-level estimator for graph classification tasks (e.g., molecular property,
67
+ bioinformatics, social graph classification).
68
+
69
+ Args:
70
+ backbone: GNN architecture (``"gin"``, ``"gcn"``, ``"gat"``, ``"sage"``,
71
+ ``"pna"``, etc.) or a custom :class:`keras.Model`. (default: ``"gin"``)
72
+ in_channels (int, optional): Size of input node features.
73
+ hidden_channels (int, optional): Dimensionality of hidden node features. (default: ``64``)
74
+ num_classes (int, optional): Number of graph classes.
75
+ num_layers (int, optional): Number of GNN layers. (default: ``3``)
76
+ pooling (str, optional): Readout pooling (``"mean"``, ``"add"``, ``"max"``). (default: ``"mean"``)
77
+ dropout (float, optional): Dropout probability. (default: ``0.5``)
78
+ **backbone_kwargs: Additional arguments forwarded to the backbone constructor.
79
+ """
80
+
81
+ def __init__(
82
+ self,
83
+ backbone: Union[str, keras.Model] = "gin",
84
+ in_channels: Optional[int] = None,
85
+ hidden_channels: int = 64,
86
+ num_classes: Optional[int] = None,
87
+ num_layers: int = 3,
88
+ pooling: str = "mean",
89
+ dropout: float = 0.5,
90
+ **backbone_kwargs,
91
+ ):
92
+ super().__init__()
93
+ self.backbone = backbone
94
+ self.in_channels = in_channels
95
+ self.hidden_channels = hidden_channels
96
+ self.num_classes = num_classes
97
+ self.num_layers = num_layers
98
+ self.pooling = pooling
99
+ self.dropout = dropout
100
+ self.backbone_kwargs = backbone_kwargs
101
+
102
+ def _init_model(self, sample_data: Any, dataset: Optional[Any] = None):
103
+ in_c = self.in_channels
104
+ if in_c is None:
105
+ if hasattr(sample_data, "num_node_features") and sample_data.num_node_features > 0:
106
+ in_c = sample_data.num_node_features
107
+ elif hasattr(sample_data, "num_features") and sample_data.num_features > 0:
108
+ in_c = sample_data.num_features
109
+ elif hasattr(sample_data, "x") and sample_data.x is not None:
110
+ in_c = int(ops.shape(sample_data.x)[-1])
111
+ else:
112
+ raise ValueError("Could not infer in_channels.")
113
+
114
+ out_c = self.num_classes
115
+ if out_c is None:
116
+ if dataset is not None and hasattr(dataset, "num_classes") and dataset.num_classes is not None:
117
+ out_c = dataset.num_classes
118
+ elif dataset is not None and isinstance(dataset, (list, tuple)):
119
+ max_y = 0
120
+ for g in dataset[:100]:
121
+ if hasattr(g, "y") and g.y is not None:
122
+ max_y = max(max_y, int(ops.convert_to_numpy(ops.max(g.y))))
123
+ out_c = max_y + 1
124
+ elif hasattr(sample_data, "num_classes") and sample_data.num_classes is not None and sample_data.num_classes > 1:
125
+ out_c = sample_data.num_classes
126
+ elif hasattr(sample_data, "y") and sample_data.y is not None:
127
+ out_c = int(ops.convert_to_numpy(ops.max(sample_data.y))) + 1
128
+ else:
129
+ out_c = 2 # default binary
130
+
131
+ out_c = max(int(out_c), 2)
132
+
133
+ self.in_channels = in_c
134
+ self.num_classes = out_c
135
+
136
+ gnn = resolve_backbone(
137
+ self.backbone,
138
+ in_channels=in_c,
139
+ out_channels=self.hidden_channels,
140
+ hidden_channels=self.hidden_channels,
141
+ num_layers=self.num_layers,
142
+ dropout=self.dropout,
143
+ **self.backbone_kwargs,
144
+ )
145
+
146
+ self.model = GraphClassificationModel(
147
+ backbone=gnn,
148
+ pooling=self.pooling,
149
+ hidden_channels=self.hidden_channels,
150
+ num_classes=out_c,
151
+ dropout=self.dropout,
152
+ )
153
+
154
+ def fit(
155
+ self,
156
+ dataset: Any,
157
+ epochs: int = 20,
158
+ lr: float = 0.01,
159
+ batch_size: int = 32,
160
+ shuffle: bool = True,
161
+ verbose: int = 1,
162
+ callbacks: Optional[List[Any]] = None,
163
+ ):
164
+ r"""Trains the graph classifier."""
165
+ if not isinstance(dataset, DataLoader):
166
+ loader = DataLoader(dataset, batch_size=batch_size, shuffle=shuffle)
167
+ sample = dataset[0]
168
+ else:
169
+ loader = dataset
170
+ sample = next(iter(loader))
171
+
172
+ if self.model is None:
173
+ self._init_model(sample, dataset=dataset)
174
+
175
+ if not self._is_compiled:
176
+ self.model.compile(
177
+ optimizer=keras.optimizers.Adam(learning_rate=lr),
178
+ loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
179
+ metrics=[keras.metrics.SparseCategoricalAccuracy(name="acc")],
180
+ )
181
+ self._is_compiled = True
182
+
183
+ history = {"loss": [], "acc": []}
184
+ for epoch in range(epochs):
185
+ batch_losses = []
186
+ batch_accs = []
187
+ for batch in loader:
188
+ x = ops.convert_to_tensor(batch.x, dtype="float32")
189
+ edge_index = ops.convert_to_tensor(batch.edge_index, dtype="int64")
190
+ batch_vec = ops.convert_to_tensor(batch.batch, dtype="int64")
191
+ y = ops.convert_to_tensor(batch.y, dtype="int64")
192
+ num_g = int(ops.shape(y)[0])
193
+ if hasattr(self.model, "num_graphs"):
194
+ self.model.num_graphs = num_g
195
+
196
+ if not self.model.built:
197
+ y_pred = self.model((x, edge_index, batch_vec), training=False)
198
+ self.model.built = True
199
+ if hasattr(self.model, "_compile_loss") and self.model._compile_loss is not None:
200
+ self.model._compile_loss.build(y, y_pred)
201
+ if hasattr(self.model, "_compile_metrics") and self.model._compile_metrics is not None:
202
+ self.model._compile_metrics.build(y, y_pred)
203
+ if self.model.optimizer is not None and not self.model.optimizer.built:
204
+ self.model.optimizer.build(self.model.trainable_variables)
205
+
206
+ res = self.model.train_on_batch((x, edge_index, batch_vec), y)
207
+ if isinstance(res, (list, tuple)):
208
+ batch_losses.append(float(res[0]))
209
+ if len(res) > 1:
210
+ batch_accs.append(float(res[1]))
211
+ else:
212
+ batch_losses.append(float(res))
213
+
214
+ avg_loss = float(np.mean(batch_losses)) if batch_losses else 0.0
215
+ avg_acc = float(np.mean(batch_accs)) if batch_accs else 0.0
216
+ history["loss"].append(avg_loss)
217
+ history["acc"].append(avg_acc)
218
+ if verbose:
219
+ print(f"Epoch {epoch + 1}/{epochs} - loss: {avg_loss:.4f} - acc: {avg_acc:.4f}")
220
+
221
+ return history
222
+
223
+ def predict_proba(self, dataset_or_loader: Any, batch_size: int = 32):
224
+ r"""Predicts class probabilities for graphs."""
225
+ if not isinstance(dataset_or_loader, DataLoader):
226
+ loader = DataLoader(dataset_or_loader, batch_size=batch_size, shuffle=False)
227
+ else:
228
+ loader = dataset_or_loader
229
+
230
+ probs = []
231
+ for batch in loader:
232
+ x = ops.convert_to_tensor(batch.x, dtype="float32")
233
+ edge_index = ops.convert_to_tensor(batch.edge_index, dtype="int64")
234
+ batch_vec = ops.convert_to_tensor(batch.batch, dtype="int64")
235
+ num_g = int(ops.convert_to_numpy(ops.max(batch_vec))) + 1 if ops.shape(batch_vec)[0] > 0 else 1
236
+ if hasattr(self.model, "num_graphs"):
237
+ self.model.num_graphs = num_g
238
+ logits = self.model((x, edge_index, batch_vec), training=False)
239
+ probs.append(ops.softmax(logits, axis=-1))
240
+ return ops.concatenate(probs, axis=0)
241
+
242
+ def predict(self, dataset_or_loader: Any, batch_size: int = 32):
243
+ r"""Predicts discrete class labels for graphs."""
244
+ probs = self.predict_proba(dataset_or_loader, batch_size=batch_size)
245
+ return ops.argmax(probs, axis=-1)
246
+
247
+ def evaluate(self, dataset_or_loader: Any, batch_size: int = 32) -> Dict[str, float]:
248
+ r"""Evaluates classification accuracy on the dataset."""
249
+ if not isinstance(dataset_or_loader, DataLoader):
250
+ loader = DataLoader(dataset_or_loader, batch_size=batch_size, shuffle=False)
251
+ else:
252
+ loader = dataset_or_loader
253
+
254
+ correct = 0
255
+ total = 0
256
+ for batch in loader:
257
+ x = ops.convert_to_tensor(batch.x, dtype="float32")
258
+ edge_index = ops.convert_to_tensor(batch.edge_index, dtype="int64")
259
+ batch_vec = ops.convert_to_tensor(batch.batch, dtype="int64")
260
+ num_g = int(ops.convert_to_numpy(ops.max(batch_vec))) + 1 if ops.shape(batch_vec)[0] > 0 else 1
261
+ if hasattr(self.model, "num_graphs"):
262
+ self.model.num_graphs = num_g
263
+ logits = self.model((x, edge_index, batch_vec), training=False)
264
+ pred = ops.argmax(logits, axis=-1)
265
+ pred_np = ops.convert_to_numpy(ops.cast(pred, "int64"))
266
+ y_np = ops.convert_to_numpy(ops.cast(batch.y, "int64"))
267
+ correct += int((pred_np == y_np).sum())
268
+ total += int(y_np.shape[0])
269
+
270
+ return {"accuracy": float(correct / max(total, 1))}