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,231 @@
1
+ from typing import Optional, Union, Callable
2
+ from keras import ops, activations
3
+ from keras.layers import Dropout
4
+
5
+ from k3_node.layers.conv.message_passing import MessagePassing
6
+ from k3_node.layers.conv.utils import gcn_norm
7
+
8
+
9
+ class ARMAConv(MessagePassing):
10
+ r"""The ARMA graph convolutional operator from the `"Graph Neural Networks
11
+ with Convolutional ARMA Filters" <https://arxiv.org/abs/1901.01343>`_
12
+ paper.
13
+
14
+ Example:
15
+ ```python
16
+ import numpy as np
17
+ from k3_node.layers import ARMAConv
18
+
19
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
20
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
21
+
22
+ layer = ARMAConv(in_channels=8, out_channels=16, num_stacks=1, num_layers=1)
23
+ out = layer(x, edge_index)
24
+ print(tuple(out.shape)) # (10, 16)
25
+ ```
26
+ """
27
+ def __init__(
28
+ self,
29
+ in_channels: int,
30
+ out_channels: Optional[int] = None,
31
+ num_stacks: int = 1,
32
+ num_layers: int = 1,
33
+ shared_weights: bool = False,
34
+ act: Union[str, Callable, None] = "relu",
35
+ dropout: float = 0.0,
36
+ bias: bool = True,
37
+ # Spektral compatibility arguments:
38
+ channels: Optional[int] = None,
39
+ order: Optional[int] = None,
40
+ iterations: Optional[int] = None,
41
+ share_weights: Optional[bool] = None,
42
+ gcn_activation: Optional[str] = None,
43
+ dropout_rate: Optional[float] = None,
44
+ activation: Optional[str] = None,
45
+ use_bias: Optional[bool] = None,
46
+ **kwargs,
47
+ ):
48
+ kwargs.setdefault("aggr", "add")
49
+ super().__init__(node_dim=0, **kwargs)
50
+
51
+ if out_channels is None:
52
+ if channels is not None:
53
+ out_channels = channels
54
+ in_channels = -1
55
+ else:
56
+ out_channels = in_channels
57
+ in_channels = -1
58
+
59
+ if order is not None:
60
+ num_stacks = order
61
+ if iterations is not None:
62
+ num_layers = iterations
63
+ if share_weights is not None:
64
+ shared_weights = share_weights
65
+ if gcn_activation is not None:
66
+ act = gcn_activation
67
+ if dropout_rate is not None:
68
+ dropout = dropout_rate
69
+ if use_bias is not None:
70
+ bias = use_bias
71
+
72
+ self.in_channels = in_channels
73
+ self.out_channels = out_channels
74
+ self.num_stacks = num_stacks
75
+ self.num_layers = num_layers
76
+ self.shared_weights = shared_weights
77
+ self.act = activations.get(act) if act is not None else None
78
+ self.dropout_rate = dropout
79
+ self.use_bias = bias
80
+
81
+ K, T, F_in, F_out = num_stacks, num_layers, in_channels, out_channels
82
+ T_w = 1 if self.shared_weights else T
83
+
84
+ self.weight = self.add_weight(
85
+ shape=(max(1, T_w - 1), K, F_out, F_out),
86
+ initializer="glorot_uniform",
87
+ name="weight",
88
+ )
89
+ if in_channels is not None and in_channels != -1:
90
+ self.init_weight = self.add_weight(
91
+ shape=(K, F_in, F_out),
92
+ initializer="glorot_uniform",
93
+ name="init_weight",
94
+ )
95
+ self.root_weight = self.add_weight(
96
+ shape=(T_w, K, F_in, F_out),
97
+ initializer="glorot_uniform",
98
+ name="root_weight",
99
+ )
100
+ else:
101
+ self.init_weight = None
102
+ self.root_weight = None
103
+
104
+ if bias:
105
+ self.bias = self.add_weight(
106
+ shape=(T_w, K, 1, F_out),
107
+ initializer="zeros",
108
+ name="bias",
109
+ )
110
+ else:
111
+ self.bias = None
112
+
113
+ self.dropout = Dropout(dropout)
114
+
115
+ def build(self, input_shape=None):
116
+ if input_shape is not None:
117
+ if isinstance(input_shape, (list, tuple)) and len(input_shape) > 0 and isinstance(input_shape[0], (list, tuple)):
118
+ dim = input_shape[0][-1]
119
+ elif isinstance(input_shape, (list, tuple)) and len(input_shape) > 0 and input_shape[0] is not None and not isinstance(input_shape[0], (int, type(None))):
120
+ dim = getattr(input_shape[0], "shape", [None, None])[-1]
121
+ else:
122
+ dim = input_shape[-1]
123
+ if (self.in_channels is None or self.in_channels == -1) and dim is not None:
124
+ self.in_channels = dim
125
+ if self.in_channels is not None and self.in_channels != -1:
126
+ K, T, F_in, F_out = self.num_stacks, self.num_layers, self.in_channels, self.out_channels
127
+ T_w = 1 if self.shared_weights else T
128
+ if self.init_weight is None:
129
+ self.init_weight = self.add_weight(
130
+ shape=(K, F_in, F_out),
131
+ initializer="glorot_uniform",
132
+ name="init_weight",
133
+ )
134
+ if self.root_weight is None:
135
+ self.root_weight = self.add_weight(
136
+ shape=(T_w, K, F_in, F_out),
137
+ initializer="glorot_uniform",
138
+ name="root_weight",
139
+ )
140
+ self.built = True
141
+
142
+ def call(self, inputs, edge_index=None, edge_weight=None, training=None, **kwargs):
143
+ if edge_index is None:
144
+ if isinstance(inputs, (list, tuple)):
145
+ if len(inputs) == 3:
146
+ x, edge_index, edge_weight = inputs
147
+ elif len(inputs) == 2:
148
+ x, edge_index = inputs
149
+ else:
150
+ raise ValueError(f"Unexpected input length {len(inputs)}")
151
+ else:
152
+ raise ValueError("Expected (x, edge_index) or x and edge_index")
153
+ else:
154
+ x = inputs
155
+
156
+ if self.in_channels is None or self.in_channels == -1:
157
+ self.build((None, ops.shape(x)[-1]))
158
+
159
+ # Legacy adj check
160
+ is_legacy = False
161
+ if hasattr(edge_index, "shape") and len(edge_index.shape) == 2:
162
+ if edge_index.shape[0] != 2 and edge_index.shape[0] == edge_index.shape[1]:
163
+ is_legacy = True
164
+ elif not hasattr(edge_index, "shape"):
165
+ is_legacy = True
166
+
167
+ if is_legacy:
168
+ if hasattr(edge_index, "indices") and not callable(edge_index.indices):
169
+ edge_weight = edge_index.values
170
+ edge_index = ops.transpose(edge_index.indices)
171
+ else:
172
+ adj = edge_index
173
+ row, col = ops.where(adj > 0)
174
+ edge_index = ops.stack([row, col], axis=0)
175
+ edge_weight = ops.take(adj, row * ops.shape(adj)[1] + col)
176
+
177
+ num_nodes = ops.shape(x)[0]
178
+ edge_index, edge_weight = gcn_norm(
179
+ edge_index,
180
+ edge_weight,
181
+ num_nodes=num_nodes,
182
+ add_self_loops=False,
183
+ dtype=x.dtype,
184
+ )
185
+
186
+ # PyG: x = x.unsqueeze(-3) -> (1, N, F_in)
187
+ # out = x
188
+ # for t in range(num_layers):
189
+ # if t == 0: out = out @ init_weight (K, F_in, F_out) -> (K, N, F_out)
190
+ # else: out = out @ weight[t-1] (K, F_out, F_out) -> (K, N, F_out)
191
+ # out = propagate(edge_index, x=out, edge_weight=edge_weight)
192
+ # root = dropout(x) @ root_weight[t] (K, F_in, F_out) -> (K, N, F_out)
193
+ # out = out + root
194
+ # if bias: out = out + bias[t]
195
+ # if act: out = act(out)
196
+ # return out.mean(dim=-3)
197
+ K, T, F_in, F_out = self.num_stacks, self.num_layers, self.in_channels, self.out_channels
198
+
199
+ # out shape: (N, K, F_out)
200
+ out = None
201
+ for t in range(self.num_layers):
202
+ w_idx = 0 if self.shared_weights else t
203
+ if t == 0:
204
+ # x: (N, F_in), init_weight: (K, F_in, F_out)
205
+ # out: (K, N, F_out)
206
+ out = ops.einsum("nf,kfo->kno", x, self.init_weight)
207
+ else:
208
+ w = self.weight[0 if self.shared_weights else t - 1]
209
+ out = ops.einsum("kno,kof->knf", out, w)
210
+
211
+ # Transpose to (N, K, F_out) so node_dim=0
212
+ out_n = ops.transpose(out, (1, 0, 2))
213
+ out_prop = self.propagate(edge_index, x=out_n, edge_weight=edge_weight, size=(num_nodes, num_nodes))
214
+ out = ops.transpose(out_prop, (1, 0, 2)) # (K, N, F_out)
215
+
216
+ root_x = self.dropout(x, training=training)
217
+ root = ops.einsum("nf,kfo->kno", root_x, self.root_weight[w_idx])
218
+ out = out + root
219
+
220
+ if self.bias is not None:
221
+ # bias[w_idx]: (K, 1, F_out)
222
+ out = out + self.bias[w_idx]
223
+
224
+ if self.act is not None:
225
+ out = self.act(out)
226
+
227
+ # out: (K, N, F_out) -> mean over K -> (N, F_out)
228
+ return ops.mean(out, axis=0)
229
+
230
+ def message(self, x_j, edge_weight=None):
231
+ return x_j if edge_weight is None else ops.expand_dims(ops.expand_dims(edge_weight, -1), -1) * x_j
@@ -0,0 +1,92 @@
1
+ from typing import Union, Tuple
2
+ from keras import layers, ops
3
+
4
+ from k3_node.layers.conv.message_passing import MessagePassing
5
+
6
+
7
+ class CGConv(MessagePassing):
8
+ r"""The Crystal Graph Convolutional operator from the
9
+ `"Crystal Graph Convolutional Neural Networks for an Accurate and
10
+ Interpretable Prediction of Material Properties"
11
+ <https://arxiv.org/abs/1710.10324>`_ paper.
12
+
13
+ Args:
14
+ channels: Size of each input sample, or a tuple for bipartite graphs.
15
+ dim: Edge feature dimensionality. (default: ``0``)
16
+ aggr: The aggregation scheme to use (``"add"``, ``"mean"``, ``"max"``).
17
+ (default: ``"add"``)
18
+ batch_norm: If set to :obj:`True`, will apply batch normalization.
19
+ (default: ``False``)
20
+ bias: If set to :obj:`False`, the layer will not learn an additive bias.
21
+ (default: ``True``)
22
+
23
+ Example:
24
+ ```python
25
+ import numpy as np
26
+ from k3_node.layers import CGConv
27
+
28
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
29
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
30
+ edge_attr = np.random.rand(30, 3).astype("float32") # 3 features per edge
31
+
32
+ layer = CGConv(channels=8, dim=3)
33
+ out = layer(x, edge_index, edge_attr)
34
+ print(tuple(out.shape)) # (10, 8)
35
+ ```
36
+ """
37
+
38
+ def __init__(
39
+ self,
40
+ channels: Union[int, Tuple[int, int]],
41
+ dim: int = 0,
42
+ aggr: str = "add",
43
+ batch_norm: bool = False,
44
+ bias: bool = True,
45
+ **kwargs,
46
+ ):
47
+ super().__init__(aggr=aggr, **kwargs)
48
+ self.channels = channels
49
+ self.dim = dim
50
+ self.batch_norm = batch_norm
51
+ self.use_bias = bias
52
+
53
+ if isinstance(channels, int):
54
+ self.in_channels_src = channels
55
+ self.in_channels_dst = channels
56
+ else:
57
+ self.in_channels_src = channels[0]
58
+ self.in_channels_dst = channels[1]
59
+
60
+ in_dim = self.in_channels_src + self.in_channels_dst + dim
61
+ self.lin_f = layers.Dense(self.in_channels_dst, use_bias=bias)
62
+ self.lin_s = layers.Dense(self.in_channels_dst, use_bias=bias)
63
+ self.bn = layers.BatchNormalization(momentum=0.9, epsilon=1e-5) if batch_norm else None
64
+
65
+ def build(self, input_shape):
66
+ in_dim = self.in_channels_src + self.in_channels_dst + self.dim
67
+ self.lin_f.build((None, in_dim))
68
+ self.lin_s.build((None, in_dim))
69
+ self.built = True
70
+
71
+ def call(self, x, edge_index=None, edge_attr=None, **kwargs):
72
+ if edge_index is None and isinstance(x, (tuple, list)):
73
+ x, edge_index = x[0], x[1]
74
+
75
+ if not isinstance(x, (tuple, list)):
76
+ x_src, x_dst = x, x
77
+ else:
78
+ x_src, x_dst = x[0], x[1]
79
+
80
+ out = self.propagate(edge_index, x=(x_src, x_dst), edge_attr=edge_attr)
81
+ if self.bn is not None:
82
+ out = self.bn(out)
83
+ out = x_dst + out
84
+ return out
85
+
86
+ def message(self, x_i, x_j, edge_attr=None):
87
+ if edge_attr is None:
88
+ z = ops.concatenate([x_i, x_j], axis=-1)
89
+ else:
90
+ z = ops.concatenate([x_i, x_j, edge_attr], axis=-1)
91
+ return ops.sigmoid(self.lin_f(z)) * ops.softplus(self.lin_s(z))
92
+
@@ -0,0 +1,137 @@
1
+ from typing import Optional, List
2
+ from keras import layers, ops
3
+
4
+ from k3_node.layers.conv.message_passing import MessagePassing
5
+ from k3_node.layers.conv.utils import get_laplacian
6
+
7
+
8
+ class ChebConv(MessagePassing):
9
+ r"""The Chebyshev spectral graph convolutional operator from the
10
+ `"Convolutional Neural Networks on Graphs with Fast Localized Spectral
11
+ Filtering" <https://arxiv.org/abs/1606.09375>`_ paper.
12
+
13
+ Args:
14
+ in_channels: Size of each input sample.
15
+ out_channels: Size of each output sample.
16
+ K: Chebyshev filter size :math:`K`.
17
+ normalization: The normalization scheme for the graph
18
+ Laplacian (``"sym"``, ``"rw"`` or :obj:`None`). (default: ``"sym"``)
19
+ bias: If set to :obj:`False`, the layer will not learn
20
+ an additive bias. (default: ``"True"``)
21
+
22
+ Example:
23
+ ```python
24
+ import numpy as np
25
+ from k3_node.layers import ChebConv
26
+
27
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
28
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
29
+
30
+ layer = ChebConv(in_channels=8, out_channels=16, K=2)
31
+ out = layer(x, edge_index)
32
+ print(tuple(out.shape)) # (10, 16)
33
+ ```
34
+ """
35
+
36
+ def __init__(
37
+ self,
38
+ in_channels: int,
39
+ out_channels: int,
40
+ K: int,
41
+ normalization: Optional[str] = "sym",
42
+ bias: bool = True,
43
+ **kwargs,
44
+ ):
45
+ super().__init__(aggr="add", **kwargs)
46
+ if K <= 0:
47
+ raise ValueError(f"K must be a positive integer, got {K}")
48
+
49
+ self.in_channels = in_channels
50
+ self.out_channels = out_channels
51
+ self.K = K
52
+ self.normalization = normalization
53
+ self.use_bias = bias
54
+
55
+ self.lins = [layers.Dense(out_channels, use_bias=False) for _ in range(K)]
56
+
57
+ def build(self, input_shape):
58
+ if isinstance(input_shape, (tuple, list)) and len(input_shape) > 0 and isinstance(input_shape[0], (tuple, list)):
59
+ feat_shape = input_shape[0]
60
+ else:
61
+ feat_shape = input_shape
62
+ for lin in self.lins:
63
+ lin.build(feat_shape)
64
+
65
+ if self.use_bias:
66
+ self.bias = self.add_weight(
67
+ shape=(self.out_channels,),
68
+ initializer="zeros",
69
+ name="bias",
70
+ )
71
+ else:
72
+ self.bias = None
73
+ self.built = True
74
+
75
+ def __norm__(
76
+ self,
77
+ edge_index,
78
+ num_nodes: Optional[int],
79
+ edge_weight=None,
80
+ normalization: Optional[str] = "sym",
81
+ lambda_max=None,
82
+ dtype=None,
83
+ ):
84
+ edge_index, edge_weight = get_laplacian(
85
+ edge_index, edge_weight, normalization, dtype, num_nodes
86
+ )
87
+ if lambda_max is None:
88
+ lambda_max = 2.0 * ops.max(edge_weight)
89
+ else:
90
+ lambda_max = ops.convert_to_tensor(lambda_max, dtype=edge_weight.dtype)
91
+
92
+ edge_weight = (2.0 * edge_weight) / lambda_max
93
+ edge_weight = ops.where(
94
+ ops.isinf(edge_weight) | ops.isnan(edge_weight), 0.0, edge_weight
95
+ )
96
+
97
+ loop_mask = edge_index[0] == edge_index[1]
98
+ edge_weight = ops.where(loop_mask, edge_weight - 1.0, edge_weight)
99
+ return edge_index, edge_weight
100
+
101
+ def call(self, x, edge_index=None, edge_weight=None, lambda_max=None, **kwargs):
102
+ if edge_index is None and isinstance(x, (tuple, list)):
103
+ x, edge_index = x[0], x[1]
104
+
105
+ num_nodes = x.shape[self.node_dim] if hasattr(x, "shape") and x.shape[self.node_dim] is not None else ops.shape(x)[self.node_dim]
106
+ edge_index, norm = self.__norm__(
107
+ edge_index,
108
+ num_nodes,
109
+ edge_weight,
110
+ self.normalization,
111
+ lambda_max,
112
+ dtype=x.dtype,
113
+ )
114
+
115
+ Tx_0 = x
116
+ Tx_1 = x
117
+ out = self.lins[0](Tx_0)
118
+
119
+ if len(self.lins) > 1:
120
+ Tx_1 = self.propagate(edge_index, x=x, norm=norm)
121
+ out = out + self.lins[1](Tx_1)
122
+
123
+ for lin in self.lins[2:]:
124
+ Tx_2 = self.propagate(edge_index, x=Tx_1, norm=norm)
125
+ Tx_2 = 2.0 * Tx_2 - Tx_0
126
+ out = out + lin(Tx_2)
127
+ Tx_0, Tx_1 = Tx_1, Tx_2
128
+
129
+ if self.bias is not None:
130
+ out = out + self.bias
131
+ return out
132
+
133
+ def message(self, x_j, norm=None):
134
+ if norm is None:
135
+ return x_j
136
+ return ops.expand_dims(norm, -1) * x_j
137
+
@@ -0,0 +1,102 @@
1
+ from keras import layers, ops
2
+
3
+ from k3_node.layers.conv.message_passing import MessagePassing
4
+ from k3_node.layers.conv.utils import (
5
+ add_self_loops,
6
+ degree,
7
+ extend_mask_for_self_loops,
8
+ remove_self_loops_masked,
9
+ )
10
+ from k3_node.ops.segment import segment_sum
11
+
12
+
13
+ class ClusterGCNConv(MessagePassing):
14
+ r"""The ClusterGCN graph convolutional operator from the
15
+ `"Cluster-GCN: An Efficient Algorithm for Training Deep and Large Graph
16
+ Convolutional Networks" <https://arxiv.org/abs/1905.07953>`_ paper.
17
+
18
+ Args:
19
+ in_channels: Size of each input sample.
20
+ out_channels: Size of each output sample.
21
+ diag_lambda: Diagonal enhancement coefficient :math:`\lambda`.
22
+ (default: ``0.0``)
23
+ add_self_loops: If set to :obj:`False`, will not add self-loops.
24
+ (default: ``True``)
25
+ bias: If set to :obj:`False`, the layer will not learn an additive bias.
26
+ (default: ``True``)
27
+
28
+ Example:
29
+ ```python
30
+ import numpy as np
31
+ from k3_node.layers import ClusterGCNConv
32
+
33
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
34
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
35
+
36
+ layer = ClusterGCNConv(in_channels=8, out_channels=16)
37
+ out = layer(x, edge_index)
38
+ print(tuple(out.shape)) # (10, 16)
39
+ ```
40
+ """
41
+
42
+ weighted_sum_message = True
43
+
44
+ def __init__(
45
+ self,
46
+ in_channels: int,
47
+ out_channels: int,
48
+ diag_lambda: float = 0.0,
49
+ add_self_loops: bool = True,
50
+ bias: bool = True,
51
+ **kwargs,
52
+ ):
53
+ super().__init__(aggr="add", **kwargs)
54
+ self.in_channels = in_channels
55
+ self.out_channels = out_channels
56
+ self.diag_lambda = diag_lambda
57
+ self.add_self_loops = add_self_loops
58
+ self.use_bias = bias
59
+
60
+ self.lin_out = layers.Dense(out_channels, use_bias=bias)
61
+ self.lin_root = layers.Dense(out_channels, use_bias=False)
62
+
63
+ def build(self, input_shape):
64
+ feat_shape = input_shape[0] if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)) else input_shape
65
+ self.lin_out.build(feat_shape)
66
+ self.lin_root.build(feat_shape)
67
+ self.built = True
68
+
69
+ def call(self, x, edge_index=None, edge_weight=None, **kwargs):
70
+ if edge_index is None and isinstance(x, (tuple, list)):
71
+ x, edge_index = x[0], x[1]
72
+
73
+ num_nodes = ops.shape(x)[self.node_dim]
74
+
75
+ keep_mask = None
76
+ if self.add_self_loops:
77
+ edge_index, _, keep_mask = remove_self_loops_masked(edge_index)
78
+ edge_index, _ = add_self_loops(edge_index, num_nodes=num_nodes)
79
+ keep_mask = extend_mask_for_self_loops(keep_mask, num_nodes)
80
+
81
+ row, col = edge_index[0], edge_index[1]
82
+ col_cast = ops.cast(col, "int32")
83
+ if keep_mask is None:
84
+ deg = degree(col_cast, num_nodes=num_nodes)
85
+ else:
86
+ deg = segment_sum(ops.cast(keep_mask, x.dtype), col_cast, num_segments=num_nodes)
87
+ deg_inv = 1.0 / ops.maximum(ops.cast(deg, x.dtype), 1.0)
88
+
89
+ edge_weight = ops.take(deg_inv, col_cast, axis=0)
90
+ loop_mask = ops.equal(row, col)
91
+ edge_weight = ops.where(loop_mask, edge_weight + self.diag_lambda * ops.take(deg_inv, col_cast, axis=0), edge_weight)
92
+ if keep_mask is not None:
93
+ edge_weight = edge_weight * ops.cast(keep_mask, edge_weight.dtype)
94
+
95
+ out = self.propagate(edge_index, x=x, edge_weight=edge_weight)
96
+ return self.lin_out(out) + self.lin_root(x)
97
+
98
+ def message(self, x_j, edge_weight=None):
99
+ if edge_weight is None:
100
+ return x_j
101
+ return ops.expand_dims(edge_weight, -1) * x_j
102
+
@@ -0,0 +1,100 @@
1
+ # ported from spektral
2
+
3
+ import warnings
4
+ from functools import wraps
5
+
6
+ from keras import ops, backend
7
+ from keras.layers import Layer
8
+
9
+ from k3_node.utils import (
10
+ is_keras_kwarg,
11
+ is_layer_kwarg,
12
+ deserialize_kwarg,
13
+ serialize_kwarg,
14
+ )
15
+
16
+
17
+ class Conv(Layer):
18
+ def __init__(self, **kwargs):
19
+ unknown = sorted(k for k in kwargs if not (is_keras_kwarg(k) or is_layer_kwarg(k)))
20
+ if unknown:
21
+ raise TypeError(f"{type(self).__name__}() got unexpected keyword argument(s): {', '.join(unknown)}")
22
+ super().__init__(**{k: v for k, v in kwargs.items() if is_keras_kwarg(k)})
23
+ self.supports_masking = True
24
+ self.kwargs_keys = []
25
+ for key in kwargs:
26
+ if is_layer_kwarg(key):
27
+ attr = kwargs[key]
28
+ attr = deserialize_kwarg(key, attr)
29
+ self.kwargs_keys.append(key)
30
+ setattr(self, key, attr)
31
+ self.call = check_dtypes_decorator(self.call)
32
+
33
+ def build(self, input_shape):
34
+ self.built = True
35
+
36
+ def call(self, inputs):
37
+ raise NotImplementedError
38
+
39
+ def get_config(self):
40
+ base_config = super().get_config()
41
+ keras_config = {}
42
+ for key in self.kwargs_keys:
43
+ keras_config[key] = serialize_kwarg(key, getattr(self, key))
44
+ return {**base_config, **keras_config, **self.config}
45
+
46
+ @property
47
+ def config(self):
48
+ return {}
49
+
50
+ @staticmethod
51
+ def preprocess(a):
52
+ return a
53
+
54
+
55
+ def check_dtypes_decorator(call):
56
+ @wraps(call)
57
+ def _inner_check_dtypes(*args, **kwargs):
58
+ if len(args) == 0:
59
+ return call(**kwargs)
60
+ elif len(args) == 1:
61
+ inputs = check_dtypes(args[0])
62
+ return call(inputs, **kwargs)
63
+ else:
64
+ checked = check_dtypes(list(args))
65
+ if isinstance(checked, (list, tuple)) and len(checked) == len(args):
66
+ return call(*checked, **kwargs)
67
+ return call(*args, **kwargs)
68
+
69
+ return _inner_check_dtypes
70
+
71
+
72
+ def check_dtypes(inputs):
73
+ if not isinstance(inputs, (list, tuple)):
74
+ return inputs
75
+ for value in inputs:
76
+ if not hasattr(value, "dtype"):
77
+ # It's not a valid tensor.
78
+ return inputs
79
+
80
+ if len(inputs) == 2:
81
+ x, a = inputs
82
+ e = None
83
+ elif len(inputs) == 3:
84
+ x, a, e = inputs
85
+ else:
86
+ return inputs
87
+
88
+ # If 'a' is an edge_index of shape (2, E), it must remain integer
89
+ if hasattr(a, "shape") and len(a.shape) == 2 and a.shape[0] == 2 and a.shape[1] != 2:
90
+ pass
91
+ elif backend.is_int_dtype(a.dtype) and backend.is_float_dtype(x.dtype):
92
+ warnings.warn(
93
+ f"The adjacency matrix of dtype {a.dtype} is incompatible with the dtype "
94
+ f"of the node features {x.dtype} and has been automatically cast to "
95
+ f"{x.dtype}."
96
+ )
97
+ a = ops.cast(a, x.dtype)
98
+
99
+ output = [_ for _ in [x, a, e] if _ is not None]
100
+ return output