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,151 @@
1
+ import numpy as np
2
+ import keras
3
+ from keras import ops
4
+ from k3_node.layers.conv.message_passing import MessagePassing
5
+ from k3_node.layers.conv.utils import scatter
6
+
7
+
8
+ class WLConv(keras.layers.Layer):
9
+ r"""The Weisfeiler Lehman (WL) operator from the `"A Reduction of a Graph
10
+ to a Canonical Form and an Algebra Arising During this Reduction"
11
+ <https://www.iti.zcu.cz/wl2018/pdf/wl_paper_translation.pdf>`_ paper.
12
+
13
+ Args:
14
+ **kwargs: Additional layer arguments.
15
+
16
+ Example:
17
+ ```python
18
+ import numpy as np
19
+ from k3_node.layers import WLConv
20
+
21
+ colors = np.array([0, 1, 0, 1]) # discrete node colors (labels)
22
+ edge_index = np.array([[0, 1, 2, 3], [1, 2, 3, 0]])
23
+ layer = WLConv()
24
+ new_colors = layer(colors, edge_index) # one Weisfeiler-Lehman refinement step
25
+ print(tuple(new_colors.shape)) # (4,)
26
+ print(tuple(layer.histogram(new_colors).shape)) # (1, 2): color histogram per graph
27
+ ```
28
+ """
29
+
30
+ def __init__(self, **kwargs):
31
+ super().__init__(**kwargs)
32
+ self.hashmap = {}
33
+
34
+ def reset_parameters(self):
35
+ self.hashmap = {}
36
+
37
+ def build(self, input_shape=None):
38
+ self.built = True
39
+
40
+ def call(self, x, edge_index=None, num_nodes=None):
41
+ if edge_index is None:
42
+ if isinstance(x, (tuple, list)) and len(x) >= 2:
43
+ x, edge_index = x[0], x[1]
44
+ elif ops.shape(x)[0] == 2:
45
+ edge_index = x
46
+ if num_nodes is None:
47
+ num_nodes = int(ops.max(edge_index)) + 1
48
+ x = ops.zeros((num_nodes,), dtype="int64")
49
+
50
+ if len(ops.shape(x)) > 1:
51
+ x = ops.argmax(x, axis=-1)
52
+
53
+ x_np = ops.convert_to_numpy(x)
54
+ edge_index_np = ops.convert_to_numpy(edge_index)
55
+
56
+ num_nodes = len(x_np)
57
+ row, col = edge_index_np[0], edge_index_np[1]
58
+
59
+ # Group neighbors by target node col
60
+ neighbors_dict = {i: [] for i in range(num_nodes)}
61
+ for src, dst in zip(row, col):
62
+ neighbors_dict[int(dst)].append(int(x_np[src]))
63
+
64
+ out = []
65
+ for i in range(num_nodes):
66
+ node_color = int(x_np[i])
67
+ sorted_neighs = sorted(neighbors_dict[i])
68
+ key = hash((node_color, tuple(sorted_neighs)))
69
+ if key not in self.hashmap:
70
+ self.hashmap[key] = len(self.hashmap)
71
+ out.append(self.hashmap[key])
72
+
73
+ return ops.convert_to_tensor(np.array(out, dtype=np.int64))
74
+
75
+ def histogram(self, x, batch=None, norm: bool = False):
76
+ x_np = ops.convert_to_numpy(x)
77
+ num_nodes = len(x_np)
78
+ if batch is None:
79
+ batch_np = np.zeros(num_nodes, dtype=np.int64)
80
+ else:
81
+ batch_np = ops.convert_to_numpy(batch).astype(np.int64)
82
+
83
+ num_colors = len(self.hashmap)
84
+ batch_size = int(np.max(batch_np)) + 1 if len(batch_np) > 0 else 1
85
+
86
+ hist = np.zeros((batch_size, num_colors), dtype=np.float32)
87
+ for b, c in zip(batch_np, x_np):
88
+ hist[b, c] += 1.0
89
+
90
+ if norm:
91
+ norms = np.linalg.norm(hist, axis=-1, keepdims=True)
92
+ hist = hist / np.maximum(norms, 1e-12)
93
+
94
+ return ops.convert_to_tensor(hist)
95
+
96
+
97
+ class WLConvContinuous(MessagePassing):
98
+ r"""The Weisfeiler Lehman operator from the `"Wasserstein
99
+ Weisfeiler-Lehman Graph Kernels" <https://arxiv.org/abs/1906.01277>`_ paper.
100
+
101
+ Args:
102
+ **kwargs: Additional arguments of :class:`MessagePassing`.
103
+
104
+ Example:
105
+ ```python
106
+ import numpy as np
107
+ from k3_node.layers import WLConvContinuous
108
+
109
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
110
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
111
+
112
+ layer = WLConvContinuous()
113
+ out = layer(x, edge_index)
114
+ print(tuple(out.shape)) # (10, 8)
115
+ ```
116
+ """
117
+
118
+ def __init__(self, **kwargs):
119
+ kwargs.setdefault("aggr", "add")
120
+ super().__init__(**kwargs)
121
+
122
+ def build(self, input_shape=None):
123
+ self.built = True
124
+
125
+ def call(self, x, edge_index, edge_weight=None):
126
+ if isinstance(x, (tuple, list)):
127
+ x_src, x_dst = x[0], x[1]
128
+ else:
129
+ x_src = x_dst = x
130
+
131
+ out = self.propagate(edge_index, x=(x_src, x_dst), edge_weight=edge_weight)
132
+
133
+ dst_index = edge_index[1]
134
+ if edge_weight is None:
135
+ edge_weight = ops.ones((ops.shape(dst_index)[0],), dtype=out.dtype)
136
+
137
+ num_nodes = ops.shape(out)[0]
138
+ deg = scatter(edge_weight, dst_index, dim=0, dim_size=num_nodes, reduce="sum")
139
+ deg_inv = ops.where(ops.equal(deg, 0), 0.0, 1.0 / deg)
140
+ out = ops.expand_dims(deg_inv, axis=-1) * out
141
+
142
+ if x_dst is not None:
143
+ out = 0.5 * (x_dst + out)
144
+
145
+ return out
146
+
147
+ def message(self, x_j, edge_weight=None):
148
+ if edge_weight is not None:
149
+ return ops.expand_dims(edge_weight, axis=-1) * x_j
150
+ return x_j
151
+
@@ -0,0 +1,187 @@
1
+ from math import ceil
2
+ from typing import Optional
3
+ import keras
4
+ from keras import ops
5
+ from k3_node.layers.conv.message_passing import MessagePassing
6
+ from k3_node.ops.creation import repeat
7
+
8
+
9
+ class GroupedConv1dFlat(keras.layers.Layer):
10
+ """A 1D convolution over input of shape (N, C, K) with kernel_size=K and groups=C."""
11
+
12
+ def __init__(self, in_channels: int, out_channels: int, kernel_size: int, **kwargs):
13
+ super().__init__(**kwargs)
14
+ self.in_channels = in_channels
15
+ self.out_channels = out_channels
16
+ self.kernel_size = kernel_size
17
+ self.multiplier = out_channels // in_channels
18
+
19
+ def build(self, input_shape=None):
20
+ self.kernel = self.add_weight(
21
+ shape=(self.in_channels, self.multiplier, self.kernel_size),
22
+ initializer="glorot_uniform",
23
+ trainable=True,
24
+ name="kernel",
25
+ )
26
+ self.bias = self.add_weight(
27
+ shape=(self.out_channels,),
28
+ initializer="zeros",
29
+ trainable=True,
30
+ name="bias",
31
+ )
32
+ super().build(input_shape)
33
+
34
+ def call(self, x):
35
+ # x: (N, C, K)
36
+ out = ops.einsum("nck,cmk->ncm", x, self.kernel)
37
+ out = ops.reshape(out, (-1, self.out_channels)) + self.bias
38
+ return out
39
+
40
+
41
+ class XConv(keras.layers.Layer):
42
+ r"""The convolutional operator on :math:`\mathcal{X}`-transformed points
43
+ from the `"PointCNN: Convolution On X-Transformed Points"
44
+ <https://arxiv.org/abs/1801.07791>`_ paper.
45
+
46
+ Args:
47
+ in_channels (int): Size of each input sample.
48
+ out_channels (int): Size of each output sample.
49
+ dim (int): Point cloud dimensionality.
50
+ kernel_size (int): Size of the convolving kernel.
51
+ hidden_channels (int, optional): Dimensionality of lifted points.
52
+ dilation (int, optional): Dilation factor. (default: :obj:`1`)
53
+ bias (bool, optional): Whether to learn an additive bias. (default: :obj:`True`)
54
+ num_workers (int, optional): Kept for PyG compatibility.
55
+
56
+ Example:
57
+ ```python
58
+ import numpy as np
59
+ from k3_node.layers import XConv
60
+
61
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
62
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
63
+ pos = np.random.rand(10, 3).astype("float32") # 3D node positions
64
+
65
+ layer = XConv(in_channels=8, out_channels=16, dim=3, kernel_size=2)
66
+ out = layer(x, pos) # neighborhoods are built from `pos`
67
+ print(tuple(out.shape)) # (10, 16)
68
+ ```
69
+ """
70
+
71
+ def __init__(
72
+ self,
73
+ in_channels: int,
74
+ out_channels: int,
75
+ dim: int,
76
+ kernel_size: int,
77
+ hidden_channels: Optional[int] = None,
78
+ dilation: int = 1,
79
+ bias: bool = True,
80
+ num_workers: int = 1,
81
+ **kwargs,
82
+ ):
83
+ super().__init__(**kwargs)
84
+
85
+ self.in_channels = in_channels
86
+ if hidden_channels is None:
87
+ hidden_channels = in_channels // 4
88
+ assert hidden_channels > 0
89
+ self.hidden_channels = hidden_channels
90
+ self.out_channels = out_channels
91
+ self.dim = dim
92
+ self.kernel_size = kernel_size
93
+ self.dilation = dilation
94
+ self.use_bias = bias
95
+
96
+ C_in, C_delta, C_out = in_channels, hidden_channels, out_channels
97
+ D, K = dim, kernel_size
98
+
99
+ # mlp1
100
+ self.mlp1_l1 = keras.layers.Dense(C_delta)
101
+ self.mlp1_bn1 = keras.layers.BatchNormalization(axis=-1, momentum=0.9, epsilon=1e-5)
102
+ self.mlp1_l2 = keras.layers.Dense(C_delta)
103
+ self.mlp1_bn2 = keras.layers.BatchNormalization(axis=-1, momentum=0.9, epsilon=1e-5)
104
+
105
+ # mlp2
106
+ self.mlp2_l1 = keras.layers.Dense(K * K)
107
+ self.mlp2_bn1 = keras.layers.BatchNormalization(axis=-1, momentum=0.9, epsilon=1e-5)
108
+ self.mlp2_conv1 = GroupedConv1dFlat(K, K * K, K)
109
+ self.mlp2_bn2 = keras.layers.BatchNormalization(axis=-1, momentum=0.9, epsilon=1e-5)
110
+ self.mlp2_conv2 = GroupedConv1dFlat(K, K * K, K)
111
+ self.mlp2_bn3 = keras.layers.BatchNormalization(axis=-1, momentum=0.9, epsilon=1e-5)
112
+
113
+ # conv
114
+ C_total = C_in + C_delta
115
+ depth_multiplier = int(ceil(C_out / C_total))
116
+ self.conv_op = GroupedConv1dFlat(C_total, C_total * depth_multiplier, K)
117
+ self.conv_lin = keras.layers.Dense(C_out, use_bias=bias)
118
+
119
+ def _mlp1(self, pos, training=None):
120
+ # pos: (N*K, D)
121
+ h = ops.elu(self.mlp1_l1(pos))
122
+ h = self.mlp1_bn1(h, training=training)
123
+ h = ops.elu(self.mlp1_l2(h))
124
+ h = self.mlp1_bn2(h, training=training)
125
+ return h
126
+
127
+ def _mlp2(self, pos_flat, training=None):
128
+ # pos_flat: (N, K * D)
129
+ K = self.kernel_size
130
+ h = ops.elu(self.mlp2_l1(pos_flat))
131
+ h = self.mlp2_bn1(h, training=training)
132
+ h = ops.reshape(h, (-1, K, K))
133
+ h = ops.elu(self.mlp2_conv1(h))
134
+ h = self.mlp2_bn2(h, training=training)
135
+ h = ops.reshape(h, (-1, K, K))
136
+ h = self.mlp2_conv2(h)
137
+ h = self.mlp2_bn3(h, training=training)
138
+ return ops.reshape(h, (-1, K, K))
139
+
140
+ def _conv(self, x_transformed):
141
+ # x_transformed: (N, C_total, K)
142
+ h = self.conv_op(x_transformed)
143
+ return self.conv_lin(h)
144
+
145
+ def call(self, x, pos, batch=None, training=None):
146
+ if len(ops.shape(pos)) == 1:
147
+ pos = ops.expand_dims(pos, axis=-1)
148
+ N = ops.shape(pos)[0]
149
+ K = self.kernel_size
150
+ D = self.dim
151
+
152
+ # Pairwise distance KNN
153
+ diff = ops.expand_dims(pos, axis=1) - ops.expand_dims(pos, axis=0)
154
+ dist = ops.sum(diff * diff, axis=-1)
155
+
156
+ if batch is not None:
157
+ mask = ops.equal(ops.expand_dims(batch, axis=1), ops.expand_dims(batch, axis=0))
158
+ dist = ops.where(mask, dist, 1e10)
159
+
160
+ _, top_k = ops.top_k(-dist, k=K * self.dilation, sorted=True)
161
+ if self.dilation > 1:
162
+ top_k = top_k[:, ::self.dilation]
163
+
164
+ row = repeat(ops.arange(N), K)
165
+ col = ops.reshape(top_k, (-1,))
166
+
167
+ pos_diff = ops.take(pos, col, axis=0) - ops.take(pos, row, axis=0)
168
+
169
+ x_star = self._mlp1(pos_diff, training=training)
170
+ x_star = ops.reshape(x_star, (N, K, self.hidden_channels))
171
+
172
+ if x is not None:
173
+ if len(ops.shape(x)) == 1:
174
+ x = ops.expand_dims(x, axis=-1)
175
+ x_col = ops.take(x, col, axis=0)
176
+ x_col = ops.reshape(x_col, (N, K, self.in_channels))
177
+ x_star = ops.concatenate([x_star, x_col], axis=-1)
178
+
179
+ x_star = ops.transpose(x_star, (0, 2, 1)) # (N, C_total, K)
180
+
181
+ transform_matrix = self._mlp2(ops.reshape(pos_diff, (N, K * D)), training=training) # (N, K, K)
182
+
183
+ x_transformed = ops.matmul(x_star, transform_matrix) # (N, C_total, K)
184
+
185
+ out = self._conv(x_transformed)
186
+ return out
187
+
@@ -0,0 +1,40 @@
1
+ r"""Dense neural network module package.
2
+
3
+ This package provides modules applicable for operating on dense tensor
4
+ representations.
5
+ """
6
+
7
+ from .linear import Linear, HeteroLinear, HeteroDictLinear
8
+ from .dense_gat_conv import DenseGATConv
9
+ from .dense_sage_conv import DenseSAGEConv
10
+ from .dense_gcn_conv import DenseGCNConv
11
+ from .dense_graph_conv import DenseGraphConv
12
+ from .dense_gin_conv import DenseGINConv
13
+ from .diff_pool import dense_diff_pool
14
+ diff_pool = dense_diff_pool
15
+ from .mincut_pool import dense_mincut_pool
16
+ mincut_pool = dense_mincut_pool
17
+ from .dmon_pool import DMoNPooling, dense_dmon_pool, dmon_pool
18
+
19
+ __all__ = [
20
+ "Linear",
21
+ "HeteroLinear",
22
+ "HeteroDictLinear",
23
+ "DenseGCNConv",
24
+ "DenseGINConv",
25
+ "DenseGraphConv",
26
+ "DenseSAGEConv",
27
+ "DenseGATConv",
28
+ "dense_diff_pool",
29
+ "diff_pool",
30
+ "dense_mincut_pool",
31
+ "mincut_pool",
32
+ "DMoNPooling",
33
+ "dense_dmon_pool",
34
+ "dmon_pool",
35
+ ]
36
+
37
+ lin_classes = __all__[:3]
38
+ conv_classes = __all__[3:8]
39
+ pool_classes = __all__[8:]
40
+
@@ -0,0 +1,149 @@
1
+ from typing import Optional
2
+ from keras import initializers, layers, ops
3
+ from .linear import Linear
4
+
5
+
6
+ class DenseGATConv(layers.Layer):
7
+ r"""See :class:`torch_geometric.nn.conv.GATConv`.
8
+
9
+ Example:
10
+ ```python
11
+ import numpy as np
12
+ from k3_node.layers import DenseGATConv
13
+
14
+ x = np.random.rand(2, 10, 8).astype("float32") # batch of 2 graphs, 10 nodes, 8 features
15
+ adj = (np.random.rand(2, 10, 10) > 0.7).astype("float32") # dense adjacency matrices
16
+
17
+ layer = DenseGATConv(in_channels=8, out_channels=16, heads=2)
18
+ out = layer(x, adj)
19
+ print(tuple(out.shape)) # (2, 10, 32)
20
+ ```
21
+ """
22
+ def __init__(
23
+ self,
24
+ in_channels: int,
25
+ out_channels: int,
26
+ heads: int = 1,
27
+ concat: bool = True,
28
+ negative_slope: float = 0.2,
29
+ dropout: float = 0.0,
30
+ bias: bool = True,
31
+ **kwargs
32
+ ):
33
+ super().__init__(**kwargs)
34
+ self.in_channels = in_channels
35
+ self.out_channels = out_channels
36
+ self.heads = heads
37
+ self.concat = concat
38
+ self.negative_slope = negative_slope
39
+ self.dropout = dropout
40
+ self.use_bias = bias
41
+
42
+ self.lin = Linear(in_channels, heads * out_channels, bias=False,
43
+ weight_initializer='glorot', name="lin")
44
+
45
+ self.att_src = self.add_weight(
46
+ shape=(1, 1, heads, out_channels),
47
+ initializer="glorot_uniform",
48
+ trainable=True,
49
+ name="att_src",
50
+ )
51
+ self.att_dst = self.add_weight(
52
+ shape=(1, 1, heads, out_channels),
53
+ initializer="glorot_uniform",
54
+ trainable=True,
55
+ name="att_dst",
56
+ )
57
+
58
+ if bias and concat:
59
+ self.bias = self.add_weight(
60
+ shape=(heads * out_channels,),
61
+ initializer="zeros",
62
+ trainable=True,
63
+ name="bias",
64
+ )
65
+ elif bias and not concat:
66
+ self.bias = self.add_weight(
67
+ shape=(out_channels,),
68
+ initializer="zeros",
69
+ trainable=True,
70
+ name="bias",
71
+ )
72
+ else:
73
+ self.bias = None
74
+
75
+ def reset_parameters(self):
76
+ self.lin.reset_parameters()
77
+ # A new initializer per tensor: a reused unseeded Keras 3 initializer returns the same values on every call.
78
+ glorot = lambda shape, dtype=None: initializers.GlorotUniform()(shape, dtype=dtype)
79
+ self.att_src.assign(glorot(self.att_src.shape, dtype=self.att_src.dtype))
80
+ self.att_dst.assign(glorot(self.att_dst.shape, dtype=self.att_dst.dtype))
81
+ if self.bias is not None:
82
+ self.bias.assign(ops.zeros(self.bias.shape, dtype=self.bias.dtype))
83
+
84
+ def build(self, input_shape):
85
+ if isinstance(input_shape, (tuple, list)):
86
+ x_shape = input_shape[0]
87
+ else:
88
+ x_shape = input_shape
89
+ self.lin.build(x_shape)
90
+ super().build(input_shape)
91
+
92
+ def call(self, x, adj, mask: Optional[any] = None, add_loop: bool = True, training=None):
93
+ is_2d_x = (len(ops.shape(x)) == 2)
94
+ if is_2d_x:
95
+ x = ops.expand_dims(x, axis=0)
96
+ if len(ops.shape(adj)) == 2:
97
+ adj = ops.expand_dims(adj, axis=0)
98
+
99
+ H, C = self.heads, self.out_channels
100
+ B = ops.shape(x)[0]
101
+ N = ops.shape(x)[1]
102
+
103
+ if add_loop:
104
+ eye = ops.expand_dims(ops.eye(N, dtype=adj.dtype), axis=0)
105
+ adj = adj * (1.0 - eye) + eye
106
+
107
+ x_proj = ops.reshape(self.lin(x), (B, N, H, C))
108
+
109
+ alpha_src = ops.sum(x_proj * self.att_src, axis=-1) # [B, N, H]
110
+ alpha_dst = ops.sum(x_proj * self.att_dst, axis=-1) # [B, N, H]
111
+
112
+ alpha = ops.expand_dims(alpha_src, axis=1) + ops.expand_dims(alpha_dst, axis=2) # [B, N, N, H]
113
+ alpha = ops.leaky_relu(alpha, negative_slope=self.negative_slope)
114
+ alpha = ops.where(ops.expand_dims(adj, axis=-1) != 0, alpha, -1e9)
115
+ alpha = ops.softmax(alpha, axis=2)
116
+
117
+ # Transpose to [B, H, N, N] and [B, H, N, C]
118
+ alpha_perm = ops.transpose(alpha, (0, 3, 1, 2))
119
+ x_perm = ops.transpose(x_proj, (0, 2, 1, 3))
120
+ out = ops.matmul(alpha_perm, x_perm) # [B, H, N, C]
121
+ out = ops.transpose(out, (0, 2, 1, 3)) # [B, N, H, C]
122
+
123
+ if self.concat:
124
+ out = ops.reshape(out, (B, N, H * C))
125
+ else:
126
+ out = ops.mean(out, axis=2)
127
+
128
+ if self.bias is not None:
129
+ out = out + self.bias
130
+
131
+ if mask is not None:
132
+ out = out * ops.cast(ops.reshape(mask, (-1, N, 1)), x.dtype)
133
+
134
+ if is_2d_x and ops.shape(out)[0] == 1:
135
+ out = ops.squeeze(out, axis=0)
136
+
137
+ return out
138
+
139
+ def compute_output_shape(self, input_shape):
140
+ if isinstance(input_shape, (tuple, list)):
141
+ x_shape = input_shape[0]
142
+ else:
143
+ x_shape = input_shape
144
+ out_dim = self.heads * self.out_channels if self.concat else self.out_channels
145
+ return (*x_shape[:-1], out_dim)
146
+
147
+ def __repr__(self) -> str:
148
+ return (f'{self.__class__.__name__}({self.in_channels}, '
149
+ f'{self.out_channels}, heads={self.heads})')
@@ -0,0 +1,117 @@
1
+ from keras import layers, ops
2
+ from .linear import Linear
3
+
4
+
5
+ class DenseGCNConv(layers.Layer):
6
+ r"""Applies the dense convolutional operator from the `"Semi-supervised
7
+ Classification with Graph Convolutional Networks"
8
+ <https://arxiv.org/abs/1609.02907>`_ paper.
9
+
10
+ .. math::
11
+ \mathbf{X}^{\prime} = \mathbf{\tilde{D}}^{-1/2} \mathbf{\tilde{A}}
12
+ \mathbf{\tilde{D}}^{-1/2} \mathbf{X} \mathbf{\Theta}
13
+
14
+ Args:
15
+ in_channels (int): Size of each input sample.
16
+ out_channels (int): Size of each output sample.
17
+ improved (bool, optional): If set to :obj:`True`, the layer computes
18
+ :math:`\mathbf{\tilde{A}} = \mathbf{A} + 2 \mathbf{I}`.
19
+ (default: :obj:`False`)
20
+ bias (bool, optional): If set to :obj:`False`, the layer will not learn
21
+ an additive bias. (default: :obj:`True`)
22
+
23
+ Example:
24
+ ```python
25
+ import numpy as np
26
+ from k3_node.layers import DenseGCNConv
27
+
28
+ x = np.random.rand(2, 10, 8).astype("float32") # batch of 2 graphs, 10 nodes, 8 features
29
+ adj = (np.random.rand(2, 10, 10) > 0.7).astype("float32") # dense adjacency matrices
30
+
31
+ layer = DenseGCNConv(in_channels=8, out_channels=16)
32
+ out = layer(x, adj)
33
+ print(tuple(out.shape)) # (2, 10, 16)
34
+ ```
35
+ """
36
+ def __init__(
37
+ self,
38
+ in_channels: int,
39
+ out_channels: int,
40
+ improved: bool = False,
41
+ bias: bool = True,
42
+ **kwargs
43
+ ):
44
+ super().__init__(**kwargs)
45
+ self.in_channels = in_channels
46
+ self.out_channels = out_channels
47
+ self.improved = improved
48
+ self.use_bias = bias
49
+
50
+ self.lin = Linear(in_channels, out_channels, bias=False,
51
+ weight_initializer='glorot', name="lin")
52
+
53
+ if bias:
54
+ self.bias = self.add_weight(
55
+ shape=(out_channels,),
56
+ initializer="zeros",
57
+ trainable=True,
58
+ name="bias",
59
+ )
60
+ else:
61
+ self.bias = None
62
+
63
+ def reset_parameters(self):
64
+ self.lin.reset_parameters()
65
+ if self.bias is not None:
66
+ self.bias.assign(ops.zeros(self.bias.shape, dtype=self.bias.dtype))
67
+
68
+ def build(self, input_shape):
69
+ if isinstance(input_shape, (tuple, list)):
70
+ x_shape = input_shape[0]
71
+ else:
72
+ x_shape = input_shape
73
+ self.lin.build(x_shape)
74
+ super().build(input_shape)
75
+
76
+ def call(self, x, adj, mask=None, add_loop: bool = True):
77
+ is_2d_x = (len(ops.shape(x)) == 2)
78
+ if is_2d_x:
79
+ x = ops.expand_dims(x, axis=0)
80
+ if len(ops.shape(adj)) == 2:
81
+ adj = ops.expand_dims(adj, axis=0)
82
+
83
+ N = ops.shape(adj)[1]
84
+
85
+ if add_loop:
86
+ diag_val = 2.0 if self.improved else 1.0
87
+ eye = ops.expand_dims(ops.eye(N, dtype=adj.dtype), axis=0)
88
+ adj = adj * (1.0 - eye) + diag_val * eye
89
+
90
+ out = self.lin(x)
91
+ deg = ops.maximum(ops.sum(adj, axis=-1), 1.0)
92
+ deg_inv_sqrt = ops.power(deg, -0.5)
93
+
94
+ adj_norm = ops.expand_dims(deg_inv_sqrt, axis=-1) * adj * ops.expand_dims(deg_inv_sqrt, axis=-2)
95
+ out = ops.matmul(adj_norm, out)
96
+
97
+ if self.bias is not None:
98
+ out = out + self.bias
99
+
100
+ if mask is not None:
101
+ out = out * ops.cast(ops.reshape(mask, (-1, N, 1)), x.dtype)
102
+
103
+ if is_2d_x and ops.shape(out)[0] == 1:
104
+ out = ops.squeeze(out, axis=0)
105
+
106
+ return out
107
+
108
+ def compute_output_shape(self, input_shape):
109
+ if isinstance(input_shape, (tuple, list)):
110
+ x_shape = input_shape[0]
111
+ else:
112
+ x_shape = input_shape
113
+ return (*x_shape[:-1], self.out_channels)
114
+
115
+ def __repr__(self) -> str:
116
+ return (f'{self.__class__.__name__}({self.in_channels}, '
117
+ f'{self.out_channels})')