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,448 @@
1
+ import math
2
+ from typing import Optional, Union, Tuple, Callable
3
+
4
+ import keras
5
+ from keras import layers, ops
6
+
7
+
8
+ def _get_activation(activation_fn: Optional[Union[str, Callable]]):
9
+ if activation_fn is None:
10
+ return None
11
+ if isinstance(activation_fn, str):
12
+ fn = activation_fn.lower()
13
+ if fn == "gelu":
14
+ return layers.Activation("gelu")
15
+ elif fn == "relu":
16
+ return layers.ReLU()
17
+ elif fn == "tanh":
18
+ return layers.Activation("tanh")
19
+ elif fn == "silu" or fn == "swish":
20
+ return layers.Activation("silu")
21
+ elif fn == "linear":
22
+ return layers.Activation("linear")
23
+ else:
24
+ return layers.Activation(activation_fn)
25
+ elif isinstance(activation_fn, layers.Layer):
26
+ return activation_fn
27
+ elif callable(activation_fn):
28
+ return layers.Activation(activation_fn)
29
+ return None
30
+
31
+
32
+ class SelfMultiheadAttentionWithPair(layers.Layer):
33
+ r"""Multihead self-attention layer supporting additive pair-level attention bias.
34
+
35
+ Args:
36
+ embed_dim (int): Total dimension of the model.
37
+ num_heads (int): Number of parallel attention heads.
38
+ dropout (float, optional): Dropout probability for attention weights. (default: ``0.1``)
39
+ bias (bool, optional): Whether to include bias terms in linear projections. (default: ``True``)
40
+ scaling_factor (float, optional): Scaling factor multiplier for attention keys. (default: ``1.0``)
41
+ **kwargs: Additional layer arguments.
42
+
43
+ Example:
44
+ ```python
45
+ import numpy as np
46
+ from k3_node.layers import SelfMultiheadAttentionWithPair
47
+
48
+ x = np.random.rand(2, 6, 16).astype("float32") # [batch, num_atoms, embed_dim]
49
+
50
+ attn_bias = np.random.rand(2 * 4, 6, 6).astype("float32") # pair bias per [batch * head]
51
+ attn = SelfMultiheadAttentionWithPair(embed_dim=16, num_heads=4)
52
+ print(tuple(attn(x, attn_bias=attn_bias).shape)) # (2, 6, 16)
53
+ ```
54
+ """
55
+
56
+ def __init__(
57
+ self,
58
+ embed_dim: int,
59
+ num_heads: int,
60
+ dropout: float = 0.1,
61
+ bias: bool = True,
62
+ scaling_factor: float = 1.0,
63
+ **kwargs,
64
+ ):
65
+ super().__init__(**kwargs)
66
+ self.embed_dim = embed_dim
67
+ self.num_heads = num_heads
68
+ self.dropout_rate = dropout
69
+ self.use_bias = bias
70
+ self.scaling_factor = float(scaling_factor)
71
+
72
+ if embed_dim % num_heads != 0:
73
+ raise ValueError(f"embed_dim ({embed_dim}) must be divisible by num_heads ({num_heads})")
74
+
75
+ self.head_dim = embed_dim // num_heads
76
+ self.scaling = (self.head_dim * self.scaling_factor) ** -0.5
77
+
78
+ self.in_proj = layers.Dense(embed_dim * 3, use_bias=bias, name="in_proj")
79
+ self.out_proj = layers.Dense(embed_dim, use_bias=bias, name="out_proj")
80
+ self.attn_dropout = layers.Dropout(dropout) if dropout > 0.0 else None
81
+
82
+ def build(self, input_shape=None):
83
+ if not self.built:
84
+ self.in_proj.build((None, None, self.embed_dim))
85
+ self.out_proj.build((None, None, self.embed_dim))
86
+ super().build(input_shape)
87
+
88
+ def call(
89
+ self,
90
+ query,
91
+ key_padding_mask=None,
92
+ attn_bias=None,
93
+ return_attn: bool = False,
94
+ training: bool = False,
95
+ ):
96
+ r"""
97
+ Args:
98
+ query (Tensor): Input tensor of shape ``[batch_size, seq_len, embed_dim]``.
99
+ key_padding_mask (Tensor, optional): Mask indicating padding tokens of shape
100
+ ``[batch_size, seq_len]`` (True/1 for padding, False/0 for valid tokens).
101
+ attn_bias (Tensor, optional): Additive pairwise attention bias of shape
102
+ ``[batch_size, num_heads, seq_len, seq_len]`` or ``[batch_size * num_heads, seq_len, seq_len]``.
103
+ return_attn (bool, optional): Whether to return raw and post-softmax attention weights.
104
+ training (bool, optional): Whether the layer is in training mode.
105
+
106
+ Returns:
107
+ Tensor or Tuple[Tensor, Tensor, Tensor]: Attention output and optionally attention weights.
108
+ """
109
+ shape = ops.shape(query)
110
+ bsz = shape[0]
111
+ tgt_len = shape[1]
112
+
113
+ qkv = self.in_proj(query) # [bsz, tgt_len, embed_dim * 3]
114
+ q, k, v = ops.split(qkv, 3, axis=-1)
115
+
116
+ # Reshape to [bsz, num_heads, seq_len, head_dim]
117
+ q = ops.reshape(q, (bsz, tgt_len, self.num_heads, self.head_dim))
118
+ q = ops.transpose(q, (0, 2, 1, 3))
119
+ q = q * self.scaling
120
+
121
+ k = ops.reshape(k, (bsz, tgt_len, self.num_heads, self.head_dim))
122
+ k = ops.transpose(k, (0, 2, 1, 3))
123
+
124
+ v = ops.reshape(v, (bsz, tgt_len, self.num_heads, self.head_dim))
125
+ v = ops.transpose(v, (0, 2, 1, 3))
126
+
127
+ # [bsz, num_heads, tgt_len, tgt_len]
128
+ attn_weights = ops.matmul(q, ops.transpose(k, (0, 1, 3, 2)))
129
+
130
+ if attn_bias is not None:
131
+ bias_shape = ops.shape(attn_bias)
132
+ if len(bias_shape) == 3:
133
+ # [bsz * num_heads, tgt_len, tgt_len] -> [bsz, num_heads, tgt_len, tgt_len]
134
+ attn_bias = ops.reshape(attn_bias, (bsz, self.num_heads, tgt_len, tgt_len))
135
+ elif len(bias_shape) == 4 and bias_shape[-1] == self.num_heads:
136
+ # [bsz, tgt_len, tgt_len, num_heads] -> [bsz, num_heads, tgt_len, tgt_len]
137
+ attn_bias = ops.transpose(attn_bias, (0, 3, 1, 2))
138
+ attn_weights = attn_weights + attn_bias
139
+
140
+ if key_padding_mask is not None:
141
+ # key_padding_mask: True or 1 for padding
142
+ mask = ops.cast(key_padding_mask, "bool")
143
+ mask = ops.expand_dims(ops.expand_dims(mask, axis=1), axis=2) # [bsz, 1, 1, tgt_len]
144
+ attn_weights = ops.where(mask, ops.cast(-1e9, attn_weights.dtype), attn_weights)
145
+
146
+ attn_probs = ops.softmax(attn_weights, axis=-1)
147
+ if self.attn_dropout is not None:
148
+ attn_probs = self.attn_dropout(attn_probs, training=training)
149
+
150
+ o = ops.matmul(attn_probs, v) # [bsz, num_heads, tgt_len, head_dim]
151
+ o = ops.transpose(o, (0, 2, 1, 3)) # [bsz, tgt_len, num_heads, head_dim]
152
+ o = ops.reshape(o, (bsz, tgt_len, self.embed_dim))
153
+ o = self.out_proj(o)
154
+
155
+ if not return_attn:
156
+ return o
157
+ return o, attn_weights, attn_probs
158
+
159
+
160
+ class TransformerEncoderLayerWithPair(layers.Layer):
161
+ r"""Transformer Encoder Layer with pair representation bias and update.
162
+
163
+ Args:
164
+ embed_dim (int, optional): Node feature dimension. (default: ``768``)
165
+ ffn_embed_dim (int, optional): Feed-forward network hidden dimension. (default: ``3072``)
166
+ attention_heads (int, optional): Number of attention heads. (default: ``8``)
167
+ dropout (float, optional): Dropout probability. (default: ``0.1``)
168
+ attention_dropout (float, optional): Attention dropout probability. (default: ``0.1``)
169
+ activation_dropout (float, optional): Activation dropout probability in FFN. (default: ``0.0``)
170
+ activation_fn (str or Callable, optional): Non-linear activation function. (default: ``"gelu"``)
171
+ post_ln (bool, optional): Whether to use Post-LN instead of Pre-LN. (default: ``False``)
172
+ **kwargs: Additional layer arguments.
173
+
174
+ Example:
175
+ ```python
176
+ import numpy as np
177
+ from k3_node.layers import TransformerEncoderLayerWithPair
178
+
179
+ x = np.random.rand(2, 6, 16).astype("float32") # [batch, num_atoms, embed_dim]
180
+
181
+ attn_bias = np.random.rand(2 * 4, 6, 6).astype("float32") # pair bias per [batch * head]
182
+ layer = TransformerEncoderLayerWithPair(embed_dim=16, ffn_embed_dim=32, attention_heads=4)
183
+ out = layer(x, attn_bias=attn_bias)
184
+ print(tuple(out[0].shape) if isinstance(out, tuple) else tuple(out.shape)) # (2, 6, 16)
185
+ ```
186
+ """
187
+
188
+ def __init__(
189
+ self,
190
+ embed_dim: int = 768,
191
+ ffn_embed_dim: int = 3072,
192
+ attention_heads: int = 8,
193
+ dropout: float = 0.1,
194
+ attention_dropout: float = 0.1,
195
+ activation_dropout: float = 0.0,
196
+ activation_fn: Union[str, Callable] = "gelu",
197
+ post_ln: bool = False,
198
+ **kwargs,
199
+ ):
200
+ super().__init__(**kwargs)
201
+ self.embed_dim = embed_dim
202
+ self.ffn_embed_dim = ffn_embed_dim
203
+ self.attention_heads = attention_heads
204
+ self.dropout_rate = dropout
205
+ self.attention_dropout_rate = attention_dropout
206
+ self.activation_dropout_rate = activation_dropout
207
+ self.activation_fn_name = activation_fn
208
+ self.post_ln = post_ln
209
+
210
+ self.self_attn = SelfMultiheadAttentionWithPair(
211
+ embed_dim=embed_dim,
212
+ num_heads=attention_heads,
213
+ dropout=attention_dropout,
214
+ name="self_attn",
215
+ )
216
+ self.self_attn_layer_norm = layers.LayerNormalization(axis=-1, epsilon=1e-5, name="self_attn_layer_norm")
217
+
218
+ self.fc1 = layers.Dense(ffn_embed_dim, name="fc1")
219
+ self.activation_fn = _get_activation(activation_fn)
220
+ self.fc2 = layers.Dense(embed_dim, name="fc2")
221
+ self.final_layer_norm = layers.LayerNormalization(axis=-1, epsilon=1e-5, name="final_layer_norm")
222
+
223
+ self.dropout1 = layers.Dropout(dropout) if dropout > 0.0 else None
224
+ self.dropout2 = layers.Dropout(dropout) if dropout > 0.0 else None
225
+ self.act_dropout = layers.Dropout(activation_dropout) if activation_dropout > 0.0 else None
226
+
227
+ def build(self, input_shape=None):
228
+ if not self.built:
229
+ self.self_attn.build((None, None, self.embed_dim))
230
+ self.self_attn_layer_norm.build((None, None, self.embed_dim))
231
+ self.fc1.build((None, None, self.embed_dim))
232
+ self.fc2.build((None, None, self.ffn_embed_dim))
233
+ self.final_layer_norm.build((None, None, self.embed_dim))
234
+ super().build(input_shape)
235
+
236
+ def call(
237
+ self,
238
+ x,
239
+ attn_bias=None,
240
+ padding_mask=None,
241
+ return_attn: bool = False,
242
+ training: bool = False,
243
+ ):
244
+ residual = x
245
+ if not self.post_ln:
246
+ x = self.self_attn_layer_norm(x)
247
+
248
+ attn_out = self.self_attn(
249
+ query=x,
250
+ key_padding_mask=padding_mask,
251
+ attn_bias=attn_bias,
252
+ return_attn=return_attn,
253
+ training=training,
254
+ )
255
+
256
+ attn_weights = None
257
+ attn_probs = None
258
+ if return_attn:
259
+ x, attn_weights, attn_probs = attn_out
260
+ else:
261
+ x = attn_out
262
+
263
+ if self.dropout1 is not None:
264
+ x = self.dropout1(x, training=training)
265
+ x = residual + x
266
+ if self.post_ln:
267
+ x = self.self_attn_layer_norm(x)
268
+
269
+ residual = x
270
+ if not self.post_ln:
271
+ x = self.final_layer_norm(x)
272
+ x = self.fc1(x)
273
+ if self.activation_fn is not None:
274
+ x = self.activation_fn(x)
275
+ if self.act_dropout is not None:
276
+ x = self.act_dropout(x, training=training)
277
+ x = self.fc2(x)
278
+ if self.dropout2 is not None:
279
+ x = self.dropout2(x, training=training)
280
+ x = residual + x
281
+ if self.post_ln:
282
+ x = self.final_layer_norm(x)
283
+
284
+ if not return_attn:
285
+ return x
286
+ return x, attn_weights, attn_probs
287
+
288
+
289
+ class TriangleMultiplication(layers.Layer):
290
+ r"""Triangle Multiplicative Update layer (AlphaFold2 / Uni-Mol2 / Uni-Mol+).
291
+
292
+ Args:
293
+ pair_dim (int): Pair feature dimension.
294
+ hidden_dim (int): Intermediate channel dimension.
295
+ mode (str, optional): Either ``"outgoing"`` or ``"incoming"``. (default: ``"outgoing"``)
296
+ **kwargs: Additional layer arguments.
297
+
298
+ Example:
299
+ ```python
300
+ import numpy as np
301
+ from k3_node.layers import TriangleMultiplication
302
+
303
+ pair = np.random.rand(2, 6, 6, 8).astype("float32") # [batch, num_atoms, num_atoms, pair_dim]
304
+ layer = TriangleMultiplication(pair_dim=8, hidden_dim=4, mode="outgoing")
305
+ print(tuple(layer(pair).shape)) # (2, 6, 6, 8)
306
+ ```
307
+ """
308
+
309
+ def __init__(
310
+ self,
311
+ pair_dim: int,
312
+ hidden_dim: int,
313
+ mode: str = "outgoing",
314
+ **kwargs,
315
+ ):
316
+ super().__init__(**kwargs)
317
+ if mode not in ("outgoing", "incoming"):
318
+ raise ValueError(f"Unknown TriangleMultiplication mode '{mode}', must be 'outgoing' or 'incoming'")
319
+ self.pair_dim = pair_dim
320
+ self.hidden_dim = hidden_dim
321
+ self.mode = mode
322
+
323
+ self.norm = layers.LayerNormalization(axis=-1, epsilon=1e-5, name="norm")
324
+ self.proj_a = layers.Dense(hidden_dim, use_bias=False, name="proj_a")
325
+ self.proj_b = layers.Dense(hidden_dim, use_bias=False, name="proj_b")
326
+ self.gate_a = layers.Dense(hidden_dim, use_bias=True, name="gate_a")
327
+ self.gate_b = layers.Dense(hidden_dim, use_bias=True, name="gate_b")
328
+
329
+ self.gate_out = layers.Dense(pair_dim, use_bias=True, name="gate_out")
330
+ self.proj_out = layers.Dense(pair_dim, use_bias=True, name="proj_out")
331
+ self.norm_out = layers.LayerNormalization(axis=-1, epsilon=1e-5, name="norm_out")
332
+
333
+ def build(self, input_shape=None):
334
+ if not self.built:
335
+ self.norm.build((None, None, None, self.pair_dim))
336
+ self.proj_a.build((None, None, None, self.pair_dim))
337
+ self.proj_b.build((None, None, None, self.pair_dim))
338
+ self.gate_a.build((None, None, None, self.pair_dim))
339
+ self.gate_b.build((None, None, None, self.pair_dim))
340
+ self.gate_out.build((None, None, None, self.pair_dim))
341
+ self.proj_out.build((None, None, None, self.hidden_dim))
342
+ self.norm_out.build((None, None, None, self.hidden_dim))
343
+ super().build(input_shape)
344
+
345
+ def call(self, pair, mask=None, training: bool = False):
346
+ r"""
347
+ Args:
348
+ pair (Tensor): Pair tensor of shape ``[batch_size, seq_len, seq_len, pair_dim]``.
349
+ mask (Tensor, optional): Optional pair mask of shape ``[batch_size, seq_len, seq_len]``.
350
+ training (bool, optional): Training flag.
351
+
352
+ Returns:
353
+ Tensor: Updated pair tensor with residual addition.
354
+ """
355
+ residual = pair
356
+ x = self.norm(pair)
357
+
358
+ a = self.proj_a(x) * ops.sigmoid(self.gate_a(x))
359
+ b = self.proj_b(x) * ops.sigmoid(self.gate_b(x))
360
+
361
+ if mask is not None:
362
+ m = ops.expand_dims(ops.cast(mask, a.dtype), axis=-1)
363
+ a = a * m
364
+ b = b * m
365
+
366
+ # Multiplicative update
367
+ if self.mode == "outgoing":
368
+ # [B, N, N, C] = \sum_k a[i, k] * b[j, k]
369
+ # transpose b to [B, K, J, C] for matrix multiply along K
370
+ # Using ops.einsum for standard formulation:
371
+ out = ops.einsum("bikc,bjkc->bijc", a, b)
372
+ else:
373
+ # incoming: \sum_k a[k, i] * b[k, j]
374
+ out = ops.einsum("bkic,bkjc->bijc", a, b)
375
+
376
+ out = self.norm_out(out)
377
+ out = self.proj_out(out) * ops.sigmoid(self.gate_out(pair))
378
+ return residual + out
379
+
380
+
381
+ class OuterProduct(layers.Layer):
382
+ r"""Outer product mean layer updating pair representations from atom representations.
383
+
384
+ Args:
385
+ embed_dim (int): Node/atom embedding dimension.
386
+ pair_dim (int): Pair representation dimension.
387
+ hidden_dim (int, optional): Intermediate projection dimension. (default: ``32``)
388
+ **kwargs: Additional layer arguments.
389
+
390
+ Example:
391
+ ```python
392
+ import numpy as np
393
+ from k3_node.layers import OuterProduct
394
+
395
+ x = np.random.rand(2, 6, 16).astype("float32") # [batch, num_atoms, embed_dim]
396
+
397
+ layer = OuterProduct(embed_dim=16, pair_dim=8, hidden_dim=4) # atom features -> pair features
398
+ print(tuple(layer(x).shape)) # (2, 6, 6, 8)
399
+ ```
400
+ """
401
+
402
+ def __init__(
403
+ self,
404
+ embed_dim: int,
405
+ pair_dim: int,
406
+ hidden_dim: int = 32,
407
+ **kwargs,
408
+ ):
409
+ super().__init__(**kwargs)
410
+ self.embed_dim = embed_dim
411
+ self.pair_dim = pair_dim
412
+ self.hidden_dim = hidden_dim
413
+
414
+ self.norm = layers.LayerNormalization(axis=-1, epsilon=1e-5, name="norm")
415
+ self.proj_left = layers.Dense(hidden_dim, use_bias=True, name="proj_left")
416
+ self.proj_right = layers.Dense(hidden_dim, use_bias=True, name="proj_right")
417
+ self.proj_out = layers.Dense(pair_dim, use_bias=True, name="proj_out")
418
+
419
+ def build(self, input_shape=None):
420
+ if not self.built:
421
+ self.norm.build((None, None, self.embed_dim))
422
+ self.proj_left.build((None, None, self.embed_dim))
423
+ self.proj_right.build((None, None, self.embed_dim))
424
+ self.proj_out.build((None, None, None, self.hidden_dim))
425
+ super().build(input_shape)
426
+
427
+ def call(self, x, mask=None, training: bool = False):
428
+ r"""
429
+ Args:
430
+ x (Tensor): Node/atom representations of shape ``[batch_size, seq_len, embed_dim]``.
431
+ mask (Tensor, optional): Node mask of shape ``[batch_size, seq_len]``.
432
+ training (bool, optional): Training flag.
433
+
434
+ Returns:
435
+ Tensor: Pair update of shape ``[batch_size, seq_len, seq_len, pair_dim]``.
436
+ """
437
+ x_norm = self.norm(x)
438
+ if mask is not None:
439
+ m = ops.expand_dims(ops.cast(mask, x.dtype), axis=-1)
440
+ x_norm = x_norm * m
441
+
442
+ left = self.proj_left(x_norm) # [B, N, C]
443
+ right = self.proj_right(x_norm) # [B, N, C]
444
+
445
+ # [B, N, 1, C] * [B, 1, N, C] -> [B, N, N, C]
446
+ prod = ops.expand_dims(left, axis=2) * ops.expand_dims(right, axis=1)
447
+ return self.proj_out(prod)
448
+
@@ -0,0 +1,187 @@
1
+ import keras
2
+ from keras import ops, layers
3
+
4
+
5
+ def _orthogonal_matrix(dim: int, seed: int = None):
6
+ # Random matrix from normal distribution
7
+ mat = keras.random.normal((dim, dim), seed=seed)
8
+ # QR decomposition to two orthogonal matrices
9
+ q, _ = ops.qr(mat, mode="reduced")
10
+ return ops.transpose(q, [1, 0])
11
+
12
+
13
+ def orthogonal_matrix(num_rows: int, num_cols: int, seed=None):
14
+ """Function ``orthogonal_matrix``.
15
+
16
+ Example:
17
+ ```python
18
+ from k3_node.layers import orthogonal_matrix
19
+
20
+ projection = orthogonal_matrix(num_rows=16, num_cols=8) # random orthogonal features (Performer)
21
+ print(tuple(projection.shape)) # (16, 8)
22
+ ```
23
+ """
24
+ num_full_blocks = int(num_rows / num_cols)
25
+ blocks = []
26
+ for _ in range(num_full_blocks):
27
+ q = _orthogonal_matrix(num_cols)
28
+ blocks.append(q)
29
+ remain_rows = num_rows - num_full_blocks * num_cols
30
+ if remain_rows > 0:
31
+ q = _orthogonal_matrix(num_cols)
32
+ blocks.append(q[:remain_rows])
33
+ mat = ops.concatenate(blocks)
34
+ return mat
35
+
36
+
37
+ def linear_attention(q, k, v):
38
+ # Plain NumPy inputs cannot be mixed with backend tensors (e.g. `ndarray @ torch.Tensor`).
39
+ """Function ``linear_attention``.
40
+
41
+ Example:
42
+ ```python
43
+ import numpy as np
44
+ from k3_node.layers import linear_attention
45
+
46
+ q = k = v = np.random.rand(1, 2, 10, 8).astype("float32") # [batch, heads, num_nodes, head_dim]
47
+ print(tuple(linear_attention(q, k, v).shape)) # (1, 2, 10, 8)
48
+ ```
49
+ """
50
+ q, k, v = ops.convert_to_tensor(q), ops.convert_to_tensor(k), ops.convert_to_tensor(v)
51
+ _k = ops.expand_dims(ops.sum(k, axis=-2), axis=-1)
52
+ D_inv = 1.0 / (q @ _k)
53
+ kv = ops.transpose(k, axes=[0, 1, 3, 2]) @ v
54
+ qkv = q @ kv
55
+ out = ops.einsum("...L,...Ld->...Ld", ops.squeeze(D_inv, axis=-1), qkv)
56
+ return out
57
+
58
+
59
+ def generalized_kernel(x, mat, kernel=ops.relu, epsilon=0.001):
60
+ """Function ``generalized_kernel``.
61
+
62
+ Example:
63
+ ```python
64
+ import numpy as np
65
+ from k3_node.layers import generalized_kernel, orthogonal_matrix
66
+
67
+ x = np.random.rand(1, 2, 10, 8).astype("float32") # [batch, heads, num_nodes, head_dim]
68
+ features = generalized_kernel(x, orthogonal_matrix(16, 8)) # random-feature map of x
69
+ print(tuple(features.shape)) # (1, 2, 10, 16)
70
+ ```
71
+ """
72
+ batch_size, num_heads = ops.shape(x)[:2]
73
+ projection = ops.transpose(mat, axes=[1, 0]) # Transpose along correct axes
74
+ projection = ops.tile(projection, [1, num_heads, 1, 1]) # Expand dimensions
75
+ x = ops.matmul(x, projection)
76
+ out = kernel(x) + epsilon
77
+ return out
78
+
79
+
80
+ class PerformerProjection(layers.Layer):
81
+ """Layer ``PerformerProjection``.
82
+
83
+ Example:
84
+ ```python
85
+ import numpy as np
86
+ from k3_node.layers import PerformerProjection
87
+
88
+ q = k = v = np.random.rand(1, 2, 10, 8).astype("float32") # [batch, heads, num_nodes, head_dim]
89
+ proj = PerformerProjection(num_cols=8) # random-feature approximation of softmax attention
90
+ print(tuple(proj(q, k, v).shape)) # (1, 2, 10, 8)
91
+ ```
92
+ """
93
+ def __init__(self, num_cols, kernel=ops.relu):
94
+ super().__init__()
95
+ import math
96
+ self.num_rows = max(1, int(num_cols * math.log(max(2, num_cols))))
97
+ self.num_cols = num_cols
98
+
99
+ # Generate an orthogonal projection matrix
100
+ self.projection_matrix = orthogonal_matrix(self.num_rows, self.num_cols)
101
+ self.kernel = kernel
102
+
103
+ def call(self, q, k, v):
104
+ q = generalized_kernel(q, self.projection_matrix, self.kernel)
105
+ k = generalized_kernel(k, self.projection_matrix, self.kernel)
106
+ out = linear_attention(q, k, v)
107
+ return out
108
+
109
+
110
+ class PerformerAttention(layers.Layer):
111
+ """
112
+ `k3_node.layers.PerformerAttention`
113
+
114
+ Initialization Arguments:
115
+
116
+ Args:
117
+ channels: The number of output channels.
118
+ heads: The number of attention heads.
119
+ head_channels: The number of attention heads.
120
+ kernel: activation function.
121
+ qkv_bias: activation function.
122
+ attn_out_bias: Bias in Attention Out.
123
+ dropout: Dropout rate.
124
+
125
+ Example:
126
+ ```python
127
+ import numpy as np
128
+ from k3_node.layers import PerformerAttention
129
+
130
+ x = np.random.rand(1, 10, 8).astype("float32") # [batch, num_nodes, channels]
131
+
132
+ mask = np.ones((1, 10), dtype=bool) # which nodes are real (not padding)
133
+ attn = PerformerAttention(channels=8, heads=2) # linear-complexity attention
134
+ print(tuple(attn(x, mask).shape)) # (1, 10, 8)
135
+ ```
136
+ """
137
+ def __init__(
138
+ self,
139
+ channels,
140
+ heads,
141
+ head_channels=64,
142
+ kernel=ops.relu,
143
+ qkv_bias=False,
144
+ attn_out_bias=True,
145
+ dropout=0.0,
146
+ ):
147
+ super().__init__()
148
+ assert channels % heads == 0
149
+ if head_channels is None:
150
+ head_channels = channels // heads
151
+
152
+ self.heads = heads
153
+ self.head_channels = head_channels
154
+ self.kernel = kernel
155
+ self.fast_attn = PerformerProjection(head_channels, kernel)
156
+
157
+ inner_channels = head_channels * heads
158
+ self.q = layers.Dense(inner_channels, use_bias=qkv_bias)
159
+ self.k = layers.Dense(inner_channels, use_bias=qkv_bias)
160
+ self.v = layers.Dense(inner_channels, use_bias=qkv_bias)
161
+ self.attn_out = layers.Dense(channels, use_bias=attn_out_bias)
162
+ self.dropout = layers.Dropout(dropout)
163
+
164
+ def call(self, x, mask=None, training=None):
165
+ B, N, *_ = x.shape
166
+ q, k, v = self.q(x), self.k(x), self.v(x)
167
+
168
+ q = ops.transpose(
169
+ ops.reshape(q, (B, N, self.heads, self.head_channels)), axes=(0, 2, 1, 3)
170
+ )
171
+ k = ops.transpose(
172
+ ops.reshape(k, (B, N, self.heads, self.head_channels)), axes=(0, 2, 1, 3)
173
+ )
174
+ v = ops.transpose(
175
+ ops.reshape(v, (B, N, self.heads, self.head_channels)), axes=(0, 2, 1, 3)
176
+ )
177
+
178
+ if mask is not None:
179
+ mask = mask[:, None, :, None]
180
+ v = ops.where(mask, v, ops.zeros_like(v))
181
+
182
+ out = self.fast_attn(q, k, v)
183
+ out = ops.transpose(out, axes=(0, 2, 1, 3)) # Transpose back
184
+ out = ops.reshape(out, (B, N, -1))
185
+ out = self.attn_out(out)
186
+ out = self.dropout(out, training=training)
187
+ return out