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,424 @@
1
+ r"""k3-node ports of `torch_geometric.nn.models`."""
2
+
3
+ from .mlp import MLP
4
+ from .attract_repel import ARLinkPredictor
5
+ from .autoencoder import InnerProductDecoder, GAE, VGAE, ARGA, ARGVA
6
+ from .deep_graph_infomax import DeepGraphInfomax
7
+ from .deepgcn import DeepGCNLayer
8
+ from .attentive_fp import AttentiveFP
9
+ from .jumping_knowledge import JumpingKnowledge, HeteroJumpingKnowledge
10
+ from .mask_label import MaskLabel
11
+ from .meta import MetaLayer
12
+ from .pmlp import PMLP
13
+ from .polynormer import Polynormer
14
+ from .basic_gnn import BasicGNN, GCN, GraphSAGE, GIN, GAT, PNA, EdgeCNN
15
+ from .label_prop import LabelPropagation
16
+ from .correct_and_smooth import CorrectAndSmooth
17
+ from .lightgcn import LightGCN, BPRLoss
18
+ from .linkx import LINKX, SparseLinear
19
+ from .rect import RECT_L
20
+ from .signed_gcn import SignedGCN
21
+ from .neural_fingerprint import NeuralFingerprint
22
+ from .graph_unet import GraphUNet
23
+ from .rev_gnn import GroupAddRev
24
+ from .sgformer import SGFormer
25
+ from .node2vec import Node2Vec
26
+ from .metapath2vec import MetaPath2Vec
27
+ from .renet import RENet
28
+ from .tgn import (
29
+ TGNMemory,
30
+ IdentityMessage,
31
+ LastAggregator,
32
+ MeanAggregator,
33
+ TimeEncoder,
34
+ LastNeighborLoader,
35
+ )
36
+ from .schnet import (
37
+ SchNet,
38
+ CFConv,
39
+ InteractionBlock as SchNetInteractionBlock,
40
+ GaussianSmearing,
41
+ ShiftedSoftplus,
42
+ RadiusInteractionGraph,
43
+ )
44
+ from .dimenet import (
45
+ DimeNet,
46
+ DimeNetPlusPlus,
47
+ BesselBasisLayer,
48
+ SphericalBasisLayer,
49
+ triplets,
50
+ )
51
+ from .gnnff import (
52
+ GNNFF,
53
+ NodeBlock,
54
+ EdgeBlock,
55
+ GaussianFilter,
56
+ )
57
+ from .gpse import (
58
+ GPSE,
59
+ GPSENodeEncoder,
60
+ GeneralLayer,
61
+ GeneralMultiLayer,
62
+ GNNStackStage,
63
+ GNNInductiveHybridMultiHead,
64
+ )
65
+ from .visnet import ViSNet
66
+ from .lpformer import LPFormer, LPAttLayer
67
+ from .graphmae2 import (
68
+ GraphMAE2,
69
+ sce_loss,
70
+ load_graphmae2_weights,
71
+ download_graphmae2_checkpoint,
72
+ )
73
+ from .graphormer import (
74
+ Graphormer,
75
+ GraphNodeFeature,
76
+ GraphAttnBias,
77
+ GraphormerMultiheadAttention,
78
+ GraphormerGraphEncoderLayer,
79
+ GraphormerGraphEncoder,
80
+ load_graphormer_weights,
81
+ download_graphormer_checkpoint,
82
+ )
83
+ from .graphormer_3d import (
84
+ Graphormer3D,
85
+ GaussianLayer,
86
+ RBF,
87
+ Graphormer3DEncoderLayer,
88
+ NodeTaskHead,
89
+ load_graphormer3d_weights,
90
+ download_graphormer3d_checkpoint,
91
+ )
92
+ from .captum import to_captum_model, to_captum_input, captum_output_to_dicts
93
+ from .gps_model import (
94
+ GPSModel,
95
+ GPSLayer,
96
+ CustomGatedGCN,
97
+ AtomEncoder,
98
+ BondEncoder,
99
+ RWSEEncoder,
100
+ SANGraphHead,
101
+ load_gps_weights,
102
+ download_gps_checkpoint,
103
+ )
104
+ from .grover import (
105
+ GROVER,
106
+ GTransEncoder,
107
+ Readout,
108
+ load_grover_weights,
109
+ download_grover_checkpoint,
110
+ )
111
+ from .mole_bert import (
112
+ MoleBERT,
113
+ MoleBERTGNN,
114
+ MoleBERTGINConv,
115
+ load_mole_bert_weights,
116
+ download_mole_bert_checkpoint,
117
+ )
118
+ from .unimol import (
119
+ UniMolModel,
120
+ UniMolConfGenModel,
121
+ UniMolDockingModel,
122
+ GaussianLayer as UniMolGaussianLayer,
123
+ NumericalEmbed as UniMolNumericalEmbed,
124
+ NonLinearHead as UniMolNonLinearHead,
125
+ DistanceHead as UniMolDistanceHead,
126
+ ClassificationHead as UniMolClassificationHead,
127
+ LinearHead as UniMolLinearHead,
128
+ MaskLMHead as UniMolMaskLMHead,
129
+ download_unimol_checkpoint,
130
+ load_unimol_weights,
131
+ )
132
+ from .unimol2 import (
133
+ UniMol2Model,
134
+ AtomFeature as UniMol2AtomFeature,
135
+ EdgeFeature as UniMol2EdgeFeature,
136
+ SE3InvariantKernel as UniMol2SE3Kernel,
137
+ MovementPredictionHead as UniMol2MovementHead,
138
+ download_unimol2_checkpoint,
139
+ load_unimol2_weights,
140
+ )
141
+ from .unimol_plus import (
142
+ UniMolPlusPCQModel,
143
+ UniMolPlusOC20Model,
144
+ EnergyHead as UniMolPlusEnergyHead,
145
+ download_unimol_plus_checkpoint,
146
+ load_unimol_plus_weights,
147
+ )
148
+ from .unimol_docking_v2 import (
149
+ DockingPoseModelV2,
150
+ download_unimol_docking_checkpoint,
151
+ load_unimol_docking_weights,
152
+ )
153
+ from . import materials
154
+ from . import bio
155
+ from . import chemistry
156
+ from .materials import (
157
+ MEGNet,
158
+ MEGNetBlock,
159
+ MEGNetGraphConv,
160
+ M3GNet,
161
+ M3GNetBlock,
162
+ M3GNetGraphConv,
163
+ ThreeBodyInteractions,
164
+ TensorNet,
165
+ TensorEmbedding,
166
+ TensorNetInteraction,
167
+ CHGNet,
168
+ CHGNetAtomGraphBlock,
169
+ CHGNetBondGraphBlock,
170
+ SO3Net,
171
+ SO3Convolution,
172
+ RealSphericalHarmonics,
173
+ GRACE,
174
+ GraceSPBasis,
175
+ GraceACEStack,
176
+ QET,
177
+ LinearQeq,
178
+ ElectrostaticPotential,
179
+ TransformedTargetModel,
180
+ Potential,
181
+ BondExpansion as MatGLBondExpansion,
182
+ GaussianExpansion as MatGLGaussianExpansion,
183
+ RadialBesselFunction as MatGLRadialBesselFunction,
184
+ FourierExpansion as MatGLFourierExpansion,
185
+ ChebyshevRadialBasis as MatGLChebyshevRadialBasis,
186
+ SphericalBesselFunction as MatGLSphericalBesselFunction,
187
+ SphericalBesselWithHarmonics as MatGLSphericalBesselWithHarmonics,
188
+ ReduceReadOut as MatGLReduceReadOut,
189
+ WeightedReadOut as MatGLWeightedReadOut,
190
+ WeightedAtomReadOut as MatGLWeightedAtomReadOut,
191
+ Set2SetReadOut as MatGLSet2SetReadOut,
192
+ EdgeSet2Set as MatGLEdgeSet2Set,
193
+ download_matgl_checkpoint,
194
+ load_matgl_weights,
195
+ load_model as load_matgl_model,
196
+ get_available_pretrained_models as get_available_matgl_models,
197
+ )
198
+
199
+ __all__ = [
200
+ "materials",
201
+ "bio",
202
+ "chemistry",
203
+ "MEGNet",
204
+ "MEGNetBlock",
205
+ "MEGNetGraphConv",
206
+ "M3GNet",
207
+ "M3GNetBlock",
208
+ "M3GNetGraphConv",
209
+ "ThreeBodyInteractions",
210
+ "TensorNet",
211
+ "TensorEmbedding",
212
+ "TensorNetInteraction",
213
+ "CHGNet",
214
+ "CHGNetAtomGraphBlock",
215
+ "CHGNetBondGraphBlock",
216
+ "SO3Net",
217
+ "SO3Convolution",
218
+ "RealSphericalHarmonics",
219
+ "GRACE",
220
+ "GraceSPBasis",
221
+ "GraceACEStack",
222
+ "QET",
223
+ "LinearQeq",
224
+ "ElectrostaticPotential",
225
+ "TransformedTargetModel",
226
+ "Potential",
227
+ "MatGLBondExpansion",
228
+ "MatGLGaussianExpansion",
229
+ "MatGLRadialBesselFunction",
230
+ "MatGLFourierExpansion",
231
+ "MatGLChebyshevRadialBasis",
232
+ "MatGLSphericalBesselFunction",
233
+ "MatGLSphericalBesselWithHarmonics",
234
+ "MatGLReduceReadOut",
235
+ "MatGLWeightedReadOut",
236
+ "MatGLWeightedAtomReadOut",
237
+ "MatGLSet2SetReadOut",
238
+ "MatGLEdgeSet2Set",
239
+ "download_matgl_checkpoint",
240
+ "load_matgl_weights",
241
+ "load_matgl_model",
242
+ "get_available_matgl_models",
243
+ "MLP",
244
+ "ARLinkPredictor",
245
+ "InnerProductDecoder",
246
+ "GAE",
247
+ "VGAE",
248
+ "ARGA",
249
+ "ARGVA",
250
+ "DeepGraphInfomax",
251
+ "DeepGCNLayer",
252
+ "AttentiveFP",
253
+ "JumpingKnowledge",
254
+ "HeteroJumpingKnowledge",
255
+ "MaskLabel",
256
+ "MetaLayer",
257
+ "PMLP",
258
+ "Polynormer",
259
+ "BasicGNN",
260
+ "GCN",
261
+ "GraphSAGE",
262
+ "GIN",
263
+ "GAT",
264
+ "PNA",
265
+ "EdgeCNN",
266
+ "LabelPropagation",
267
+ "CorrectAndSmooth",
268
+ "LightGCN",
269
+ "BPRLoss",
270
+ "LINKX",
271
+ "SparseLinear",
272
+ "RECT_L",
273
+ "SignedGCN",
274
+ "NeuralFingerprint",
275
+ "GraphUNet",
276
+ "GroupAddRev",
277
+ "SGFormer",
278
+ "Node2Vec",
279
+ "MetaPath2Vec",
280
+ "RENet",
281
+ "TGNMemory",
282
+ "IdentityMessage",
283
+ "LastAggregator",
284
+ "MeanAggregator",
285
+ "TimeEncoder",
286
+ "LastNeighborLoader",
287
+ "SchNet",
288
+ "CFConv",
289
+ "SchNetInteractionBlock",
290
+ "GaussianSmearing",
291
+ "ShiftedSoftplus",
292
+ "RadiusInteractionGraph",
293
+ "DimeNet",
294
+ "DimeNetPlusPlus",
295
+ "BesselBasisLayer",
296
+ "SphericalBasisLayer",
297
+ "triplets",
298
+ "GNNFF",
299
+ "NodeBlock",
300
+ "EdgeBlock",
301
+ "GaussianFilter",
302
+ "GPSE",
303
+ "GPSENodeEncoder",
304
+ "GeneralLayer",
305
+ "GeneralMultiLayer",
306
+ "GNNStackStage",
307
+ "GNNInductiveHybridMultiHead",
308
+ "ViSNet",
309
+ "LPFormer",
310
+ "LPAttLayer",
311
+ "GraphMAE2",
312
+ "sce_loss",
313
+ "load_graphmae2_weights",
314
+ "download_graphmae2_checkpoint",
315
+ "Graphormer",
316
+ "GraphNodeFeature",
317
+ "GraphAttnBias",
318
+ "GraphormerMultiheadAttention",
319
+ "GraphormerGraphEncoderLayer",
320
+ "GraphormerGraphEncoder",
321
+ "load_graphormer_weights",
322
+ "download_graphormer_checkpoint",
323
+ "Graphormer3D",
324
+ "GaussianLayer",
325
+ "RBF",
326
+ "Graphormer3DEncoderLayer",
327
+ "NodeTaskHead",
328
+ "load_graphormer3d_weights",
329
+ "download_graphormer3d_checkpoint",
330
+ "to_captum_model",
331
+ "to_captum_input",
332
+ "captum_output_to_dicts",
333
+ "GPSModel",
334
+ "GPSLayer",
335
+ "CustomGatedGCN",
336
+ "AtomEncoder",
337
+ "BondEncoder",
338
+ "RWSEEncoder",
339
+ "SANGraphHead",
340
+ "load_gps_weights",
341
+ "download_gps_checkpoint",
342
+ "GROVER",
343
+ "GTransEncoder",
344
+ "Readout",
345
+ "load_grover_weights",
346
+ "download_grover_checkpoint",
347
+ "MoleBERT",
348
+ "MoleBERTGNN",
349
+ "MoleBERTGINConv",
350
+ "load_mole_bert_weights",
351
+ "download_mole_bert_checkpoint",
352
+ "UniMolModel",
353
+ "UniMolConfGenModel",
354
+ "UniMolDockingModel",
355
+ "UniMolGaussianLayer",
356
+ "UniMolNumericalEmbed",
357
+ "UniMolNonLinearHead",
358
+ "UniMolDistanceHead",
359
+ "UniMolClassificationHead",
360
+ "UniMolLinearHead",
361
+ "UniMolMaskLMHead",
362
+ "download_unimol_checkpoint",
363
+ "load_unimol_weights",
364
+ "UniMol2Model",
365
+ "UniMol2AtomFeature",
366
+ "UniMol2EdgeFeature",
367
+ "UniMol2SE3Kernel",
368
+ "UniMol2MovementHead",
369
+ "download_unimol2_checkpoint",
370
+ "load_unimol2_weights",
371
+ "UniMolPlusPCQModel",
372
+ "UniMolPlusOC20Model",
373
+ "UniMolPlusEnergyHead",
374
+ "download_unimol_plus_checkpoint",
375
+ "load_unimol_plus_weights",
376
+ "DockingPoseModelV2",
377
+ "download_unimol_docking_checkpoint",
378
+ "load_unimol_docking_weights",
379
+ ]
380
+
381
+ # Inject Hugging Face Hub capabilities (from_pretrained, save_pretrained, push_to_hub, predict)
382
+ # to all models in k3_node.models
383
+ import keras
384
+ from k3_node.hub.hub_mixin import K3NodeHubMixin
385
+
386
+ def _defines_own(cls, attr):
387
+ """True if a k3_node class in ``cls``'s MRO defines ``attr`` itself (e.g. a model-specific
388
+ ``from_pretrained`` that loads original checkpoints), which must not be overwritten."""
389
+ return any(attr in vars(klass) for klass in cls.__mro__ if klass.__module__.startswith("k3_node"))
390
+
391
+
392
+ def _is_graph_input(data):
393
+ if hasattr(data, "edge_index") or (hasattr(data, "z") and hasattr(data, "pos")):
394
+ return True
395
+ return isinstance(data, dict) and any(k in data for k in ("edge_index", "pos", "z"))
396
+
397
+
398
+ def _graph_aware_predict(self, data=None, *args, **kwargs):
399
+ """Graph inputs (``Data``, ``Batch``, graph dicts) use the hub-style ``predict``; everything
400
+ else keeps Keras' batched ``Model.predict`` (arrays, ``tf.data``, ``PyDataset``, ...)."""
401
+ if _is_graph_input(data):
402
+ return K3NodeHubMixin.predict(self, data, *args, **kwargs)
403
+ return keras.Model.predict(self, data, *args, **kwargs)
404
+
405
+
406
+ for _name in list(__all__):
407
+ _obj = globals().get(_name)
408
+ if isinstance(_obj, type) and issubclass(_obj, (keras.Model, keras.layers.Layer)):
409
+ if not issubclass(_obj, K3NodeHubMixin):
410
+ if not _defines_own(_obj, "from_pretrained"):
411
+ _obj.from_pretrained = classmethod(K3NodeHubMixin.from_pretrained.__func__)
412
+ if not _defines_own(_obj, "save_pretrained"):
413
+ _obj.save_pretrained = K3NodeHubMixin.save_pretrained
414
+ if not _defines_own(_obj, "push_to_hub"):
415
+ _obj.push_to_hub = K3NodeHubMixin.push_to_hub
416
+ if not _defines_own(_obj, "predict"):
417
+ if issubclass(_obj, keras.Model):
418
+ _obj.predict = _graph_aware_predict
419
+ else:
420
+ _obj.predict = K3NodeHubMixin.predict
421
+ if not hasattr(_obj, "_get_config"):
422
+ _obj._get_config = K3NodeHubMixin._get_config
423
+
424
+
@@ -0,0 +1,232 @@
1
+ from typing import Optional
2
+ import keras
3
+ from keras import ops
4
+
5
+ from k3_node.layers.conv import GATConv
6
+ from k3_node.layers.conv.message_passing import MessagePassing
7
+ from k3_node.layers.conv.utils import softmax
8
+ from k3_node.layers.pool import global_add_pool
9
+ from k3_node.layers.pool.glob import _infer_size
10
+ from k3_node.ops.segment import segment_sum
11
+
12
+ try:
13
+ from keras.src.backend.common.symbolic_scope import in_symbolic_scope
14
+ except ImportError:
15
+ def in_symbolic_scope():
16
+ return False
17
+
18
+
19
+ def _gru_step(cell, x, h):
20
+ out, _ = cell(x, [h])
21
+ return out
22
+
23
+
24
+ class GATEConv(MessagePassing):
25
+ r"""The edge-conditioned attention layer used as the first message
26
+ passing step of `AttentiveFP`."""
27
+ def __init__(self, in_channels: int, out_channels: int, edge_dim: int,
28
+ dropout: float = 0.0, **kwargs):
29
+ super().__init__(aggr="add", node_dim=0, **kwargs)
30
+
31
+ self.in_channels = in_channels
32
+ self.out_channels = out_channels
33
+ self.edge_dim = edge_dim
34
+ self.dropout_rate = dropout
35
+
36
+ self.lin1 = keras.layers.Dense(out_channels, use_bias=False)
37
+ self.lin2 = keras.layers.Dense(out_channels, use_bias=False)
38
+ self.dropout = keras.layers.Dropout(dropout) if dropout > 0.0 else None
39
+
40
+ self.lin1.build((None, in_channels + edge_dim))
41
+ self.lin2.build((None, out_channels))
42
+
43
+ self.att_l = self.add_weight(shape=(1, out_channels), initializer="glorot_uniform", name="att_l")
44
+ self.att_r = self.add_weight(shape=(1, in_channels), initializer="glorot_uniform", name="att_r")
45
+ self.bias = self.add_weight(shape=(out_channels,), initializer="zeros", name="bias")
46
+
47
+ def build(self, input_shape=None):
48
+ self.built = True
49
+
50
+ def call(self, x, edge_index, edge_attr, training=None):
51
+ row, col = ops.cast(edge_index[0], "int32"), ops.cast(edge_index[1], "int32")
52
+
53
+ x_j = ops.take(x, row, axis=0)
54
+ x_i = ops.take(x, col, axis=0)
55
+
56
+ edge_attr = ops.cast(edge_attr, x.dtype)
57
+ h = ops.leaky_relu(self.lin1(ops.concatenate([x_j, edge_attr], axis=-1)), negative_slope=0.01)
58
+ alpha_j = ops.sum(h * self.att_l, axis=-1)
59
+ alpha_i = ops.sum(x_i * self.att_r, axis=-1)
60
+ alpha = ops.leaky_relu(alpha_j + alpha_i, negative_slope=0.01)
61
+
62
+ num_nodes = ops.shape(x)[0]
63
+ alpha = softmax(alpha, col, num_nodes=num_nodes, dim=0)
64
+ if self.dropout is not None:
65
+ alpha = self.dropout(alpha, training=training)
66
+
67
+ message = self.lin2(x_j) * ops.expand_dims(alpha, -1)
68
+ out = segment_sum(message, col, num_segments=num_nodes)
69
+ return out + self.bias
70
+
71
+
72
+ class AttentiveFP(keras.Model):
73
+ r"""The Attentive FP model for molecular representation learning from the
74
+ `"Pushing the Boundaries of Molecular Representation for Drug Discovery
75
+ with the Graph Attention Mechanism"
76
+ <https://pubs.acs.org/doi/10.1021/acs.jmedchem.9b00959>`_ paper, based on
77
+ graph attention mechanisms.
78
+
79
+ Args:
80
+ in_channels (int): Size of each input sample.
81
+ hidden_channels (int): Hidden node feature dimensionality.
82
+ out_channels (int): Size of each output sample.
83
+ edge_dim (int): Edge feature dimensionality.
84
+ num_layers (int): Number of GNN layers.
85
+ num_timesteps (int): Number of iterative refinement steps for global
86
+ readout.
87
+ dropout (float, optional): Dropout probability. (default: `0.0`)
88
+ batch_size (int, optional): Fixed batch size (number of graphs) for JAX/XLA
89
+ static shape compatibility. (default: `None`)
90
+
91
+ Example:
92
+ ```python
93
+ import numpy as np
94
+ from k3_node.models import AttentiveFP
95
+
96
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
97
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
98
+ edge_attr = np.random.rand(30, 3).astype("float32") # bond features
99
+
100
+ batch = np.repeat([0, 1], 5) # two molecules with 5 atoms each
101
+ model = AttentiveFP(in_channels=8, hidden_channels=16, out_channels=1, edge_dim=3,
102
+ num_layers=2, num_timesteps=2)
103
+ out = model(x, edge_index, edge_attr, batch) # one prediction per molecule
104
+ print(tuple(out.shape)) # (2, 1)
105
+ ```
106
+ """
107
+ def __init__(
108
+ self,
109
+ in_channels: int,
110
+ hidden_channels: int,
111
+ out_channels: int,
112
+ edge_dim: int,
113
+ num_layers: int,
114
+ num_timesteps: int,
115
+ dropout: float = 0.0,
116
+ batch_size: Optional[int] = None,
117
+ **kwargs,
118
+ ):
119
+ super().__init__(**kwargs)
120
+
121
+ self.in_channels = in_channels
122
+ self.hidden_channels = hidden_channels
123
+ self.out_channels = out_channels
124
+ self.edge_dim = edge_dim
125
+ self.num_layers = num_layers
126
+ self.num_timesteps = num_timesteps
127
+ self.dropout_rate = dropout
128
+ self.batch_size = batch_size
129
+
130
+ self.lin1 = keras.layers.Dense(hidden_channels)
131
+ self.lin1.build((None, in_channels))
132
+
133
+ self.gate_conv = GATEConv(hidden_channels, hidden_channels, edge_dim, dropout)
134
+ self.gru = keras.layers.GRUCell(hidden_channels)
135
+ self.gru.build((None, hidden_channels))
136
+
137
+ self.atom_convs = []
138
+ self.atom_grus = []
139
+ for _ in range(num_layers - 1):
140
+ conv = GATConv(hidden_channels, hidden_channels, dropout=dropout,
141
+ add_self_loops=False, negative_slope=0.01)
142
+ conv.build((None, hidden_channels))
143
+ self.atom_convs.append(conv)
144
+ gru = keras.layers.GRUCell(hidden_channels)
145
+ gru.build((None, hidden_channels))
146
+ self.atom_grus.append(gru)
147
+
148
+ self.mol_conv = GATConv(hidden_channels, hidden_channels, dropout=dropout,
149
+ add_self_loops=False, negative_slope=0.01)
150
+ self.mol_conv.build([(None, hidden_channels), (None, hidden_channels)])
151
+ self.mol_gru = keras.layers.GRUCell(hidden_channels)
152
+ self.mol_gru.build((None, hidden_channels))
153
+
154
+ self.lin2 = keras.layers.Dense(out_channels)
155
+ self.lin2.build((None, hidden_channels))
156
+
157
+ self.dropout = keras.layers.Dropout(dropout) if dropout > 0.0 else None
158
+ self.built = True
159
+
160
+ def call(self, x, edge_index=None, edge_attr=None, batch=None, batch_size=None, training=None):
161
+ if isinstance(x, dict):
162
+ edge_index = x.get("edge_index")
163
+ edge_attr = x.get("edge_attr")
164
+ batch = x.get("batch")
165
+ batch_size = x.get("batch_size", batch_size)
166
+ x = x.get("x")
167
+ elif isinstance(x, (tuple, list)) and edge_index is None:
168
+ if len(x) >= 4:
169
+ x, edge_index, edge_attr, batch = x[0], x[1], x[2], x[3]
170
+ elif len(x) == 3:
171
+ x, edge_index, edge_attr = x[0], x[1], x[2]
172
+
173
+ bs = batch_size if batch_size is not None else self.batch_size
174
+ x = ops.cast(x, "float32")
175
+ if edge_attr is not None:
176
+ edge_attr = ops.cast(edge_attr, "float32")
177
+ # Atom Embedding:
178
+ x = ops.leaky_relu(self.lin1(x), negative_slope=0.01)
179
+
180
+ h = ops.elu(self.gate_conv(x, edge_index, edge_attr, training=training))
181
+ if self.dropout is not None:
182
+ h = self.dropout(h, training=training)
183
+ x = ops.relu(_gru_step(self.gru, h, x))
184
+
185
+ for conv, gru in zip(self.atom_convs, self.atom_grus):
186
+ h = conv(x, edge_index, training=training)
187
+ h = ops.elu(h)
188
+ if self.dropout is not None:
189
+ h = self.dropout(h, training=training)
190
+ x = ops.relu(_gru_step(gru, h, x))
191
+
192
+ # Molecule Embedding:
193
+ if batch is None:
194
+ batch = ops.zeros((ops.shape(x)[0],), dtype="int32")
195
+ else:
196
+ batch = ops.cast(batch, "int32")
197
+ num_nodes = ops.shape(batch)[0]
198
+ row = ops.arange(num_nodes, dtype="int32")
199
+ mol_edge_index = ops.stack([row, batch], axis=0)
200
+
201
+ if keras.config.backend() == "jax":
202
+ size = bs
203
+ elif in_symbolic_scope():
204
+ size = bs
205
+ else:
206
+ size = _infer_size(batch)
207
+ if size is None:
208
+ size = bs
209
+
210
+ out = ops.relu(global_add_pool(x, batch, size=size))
211
+ for _ in range(self.num_timesteps):
212
+ h = ops.elu(self.mol_conv((x, out), mol_edge_index, training=training))
213
+ if self.dropout is not None:
214
+ h = self.dropout(h, training=training)
215
+ out = ops.relu(_gru_step(self.mol_gru, h, out))
216
+
217
+ # Predictor:
218
+ if self.dropout is not None:
219
+ out = self.dropout(out, training=training)
220
+ return self.lin2(out)
221
+
222
+ def __repr__(self) -> str:
223
+ return (
224
+ f"{self.__class__.__name__}("
225
+ f"in_channels={self.in_channels}, "
226
+ f"hidden_channels={self.hidden_channels}, "
227
+ f"out_channels={self.out_channels}, "
228
+ f"edge_dim={self.edge_dim}, "
229
+ f"num_layers={self.num_layers}, "
230
+ f"num_timesteps={self.num_timesteps}"
231
+ f")"
232
+ )