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,175 @@
1
+ from typing import Dict, List, Optional, Tuple, Union
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 softmax
6
+
7
+
8
+ def semantic_group(xs: List[any], q: any, k_lin: keras.layers.Layer) -> Tuple[Optional[any], Optional[any]]:
9
+ if len(xs) == 0:
10
+ return None, None
11
+ num_edge_types = len(xs)
12
+ out = ops.stack(xs, axis=0) # (E, N, C)
13
+ if ops.size(out) == 0:
14
+ return ops.reshape(out, (0, ops.shape(out)[-1])), None
15
+
16
+ k_proj = ops.tanh(k_lin(out)) # (E, N, C)
17
+ mean_k = ops.mean(k_proj, axis=1) # (E, C)
18
+ attn_score = ops.sum(q * mean_k, axis=-1) # (E,)
19
+ if num_edge_types == 1:
20
+ attn = ops.ones_like(attn_score)
21
+ else:
22
+ attn = ops.softmax(attn_score, axis=0) # (E,)
23
+
24
+ attn_expanded = ops.reshape(attn, (num_edge_types, 1, 1))
25
+ out = ops.sum(attn_expanded * out, axis=0) # (N, C)
26
+ return out, attn
27
+
28
+
29
+ class HANConv(MessagePassing):
30
+ r"""The Heterogeneous Graph Attention Operator from the
31
+ `"Heterogeneous Graph Attention Network" <https://arxiv.org/abs/1903.07293>`_ paper.
32
+
33
+ Args:
34
+ in_channels (int or Dict[str, int]): Size of each input sample of every node type.
35
+ out_channels (int): Size of each output sample.
36
+ metadata (Tuple[List[str], List[Tuple[str, str, str]]]): Node types and edge types.
37
+ heads (int, optional): Number of multi-head-attentions. (default: :obj:`1`)
38
+ negative_slope (float, optional): LeakyReLU angle of the negative slope. (default: :obj:`0.2`)
39
+
40
+ Example:
41
+ ```python
42
+ import numpy as np
43
+ from k3_node.layers import HANConv
44
+
45
+ x_dict = {
46
+ "author": np.random.rand(3, 8).astype("float32"),
47
+ "paper": np.random.rand(4, 8).astype("float32"),
48
+ }
49
+ edge_index_dict = {
50
+ ("author", "writes", "paper"): np.array([[0, 1, 2], [1, 2, 3]]),
51
+ ("paper", "written_by", "author"): np.array([[1, 2, 3], [0, 1, 2]]),
52
+ }
53
+ metadata = (list(x_dict), list(edge_index_dict)) # (node types, edge types)
54
+
55
+ layer = HANConv(in_channels={"author": 8, "paper": 8}, out_channels=16, metadata=metadata, heads=2)
56
+ out_dict = layer(x_dict, edge_index_dict)
57
+ print(tuple(out_dict["author"].shape), tuple(out_dict["paper"].shape)) # (3, 16) (4, 16)
58
+ ```
59
+ """
60
+
61
+ def __init__(
62
+ self,
63
+ in_channels: Union[int, Dict[str, int]],
64
+ out_channels: int,
65
+ metadata: Tuple[List[str], List[Tuple[str, str, str]]],
66
+ heads: int = 1,
67
+ negative_slope: float = 0.2,
68
+ **kwargs,
69
+ ):
70
+ super().__init__(aggr="add", node_dim=0, **kwargs)
71
+
72
+ if not isinstance(in_channels, dict):
73
+ in_channels = {node_type: in_channels for node_type in metadata[0]}
74
+
75
+ self.heads = heads
76
+ self.in_channels = in_channels
77
+ self.out_channels = out_channels
78
+ self.negative_slope = negative_slope
79
+ self.metadata = metadata
80
+
81
+ self.k_lin = keras.layers.Dense(out_channels, use_bias=True)
82
+ self.proj = {
83
+ node_type: keras.layers.Dense(out_channels, use_bias=True)
84
+ for node_type, ch in self.in_channels.items()
85
+ }
86
+
87
+ def build(self, input_shape=None):
88
+ H, D = self.heads, self.out_channels // self.heads
89
+
90
+ self.q = self.add_weight(
91
+ shape=(1, self.out_channels),
92
+ initializer="glorot_uniform",
93
+ trainable=True,
94
+ name="q",
95
+ )
96
+
97
+ self.lin_src = {}
98
+ self.lin_dst = {}
99
+ for edge_type in self.metadata[1]:
100
+ key = "__".join(edge_type)
101
+ self.lin_src[key] = self.add_weight(
102
+ shape=(1, H, D),
103
+ initializer="glorot_uniform",
104
+ trainable=True,
105
+ name=f"lin_src_{key}",
106
+ )
107
+ self.lin_dst[key] = self.add_weight(
108
+ shape=(1, H, D),
109
+ initializer="glorot_uniform",
110
+ trainable=True,
111
+ name=f"lin_dst_{key}",
112
+ )
113
+
114
+ for nt, layer in self.proj.items():
115
+ if not layer.built:
116
+ layer.build((None, self.in_channels[nt]))
117
+ if not self.k_lin.built:
118
+ self.k_lin.build((None, self.out_channels))
119
+
120
+ super().build(input_shape)
121
+
122
+ def call(
123
+ self,
124
+ x_dict: Dict[str, any],
125
+ edge_index_dict: Dict[Tuple[str, str, str], any],
126
+ return_semantic_attention_weights: bool = False,
127
+ ):
128
+ if not self.built:
129
+ self.build()
130
+
131
+ H, D = self.heads, self.out_channels // self.heads
132
+ x_node_dict = {}
133
+ out_dict = {nt: [] for nt in self.metadata[0]}
134
+
135
+ for node_type, x in x_dict.items():
136
+ proj_x = self.proj[node_type](x)
137
+ x_node_dict[node_type] = ops.reshape(proj_x, (-1, H, D))
138
+
139
+ for edge_type, edge_index in edge_index_dict.items():
140
+ src_type, _, dst_type = edge_type
141
+ key = "__".join(edge_type)
142
+ lin_src = self.lin_src[key]
143
+ lin_dst = self.lin_dst[key]
144
+ x_src = x_node_dict[src_type]
145
+ x_dst = x_node_dict[dst_type]
146
+
147
+ alpha_src = ops.sum(x_src * lin_src, axis=-1) # (N_src, H)
148
+ alpha_dst = ops.sum(x_dst * lin_dst, axis=-1) # (N_dst, H)
149
+
150
+ out = self.propagate(
151
+ edge_index,
152
+ x=(x_src, x_dst),
153
+ alpha=(alpha_src, alpha_dst),
154
+ )
155
+ out = ops.relu(out)
156
+ out_dict[dst_type].append(out)
157
+
158
+ result = {}
159
+ semantic_attn_dict = {}
160
+ for node_type, outs in out_dict.items():
161
+ out, attn = semantic_group(outs, self.q, self.k_lin)
162
+ result[node_type] = out
163
+ semantic_attn_dict[node_type] = attn
164
+
165
+ if return_semantic_attention_weights:
166
+ return result, semantic_attn_dict
167
+ return result
168
+
169
+ def message(self, x_j, alpha_i, alpha_j, index=None):
170
+ alpha = alpha_j + alpha_i
171
+ alpha = ops.leaky_relu(alpha, negative_slope=self.negative_slope)
172
+ alpha = softmax(alpha, index=index)
173
+ out = x_j * ops.expand_dims(alpha, axis=-1)
174
+ return ops.reshape(out, (-1, self.out_channels))
175
+
@@ -0,0 +1,131 @@
1
+ import keras
2
+ from keras import ops
3
+ from k3_node.layers.conv.message_passing import MessagePassing
4
+ from k3_node.layers.conv.utils import softmax
5
+ from k3_node.layers.dense.linear import HeteroLinear
6
+
7
+
8
+ class HEATConv(MessagePassing):
9
+ r"""The heterogeneous edge-enhanced graph attentional operator from the
10
+ `"Heterogeneous Edge-Enhanced Graph Attention Network For Multi-Agent
11
+ Trajectory Prediction" <https://arxiv.org/abs/2106.07161>`_ paper.
12
+
13
+ Args:
14
+ in_channels (int): Size of each input sample.
15
+ out_channels (int): Size of each output sample.
16
+ num_node_types (int): The number of node types.
17
+ num_edge_types (int): The number of edge types.
18
+ edge_type_emb_dim (int): The embedding size of edge types.
19
+ edge_dim (int): Edge feature dimensionality.
20
+ edge_attr_emb_dim (int): The embedding size of edge features.
21
+ heads (int, optional): Number of multi-head-attentions. (default: :obj:`1`)
22
+ concat (bool, optional): Whether to concatenate multi-head attention. (default: :obj:`True`)
23
+ negative_slope (float, optional): LeakyReLU angle. (default: :obj:`0.2`)
24
+ root_weight (bool, optional): Whether to add root node features. (default: :obj:`True`)
25
+ bias (bool, optional): Whether to learn an additive bias. (default: :obj:`True`)
26
+
27
+ Example:
28
+ ```python
29
+ import numpy as np
30
+ from k3_node.layers import HEATConv
31
+
32
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
33
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
34
+
35
+ node_type = np.random.randint(0, 2, size=(10,))
36
+ edge_type = np.random.randint(0, 3, size=(30,))
37
+ edge_attr = np.random.rand(30, 5).astype("float32")
38
+ layer = HEATConv(in_channels=8, out_channels=16, num_node_types=2, num_edge_types=3,
39
+ edge_type_emb_dim=4, edge_dim=5, edge_attr_emb_dim=6, heads=2)
40
+ out = layer(x, edge_index, node_type, edge_type, edge_attr)
41
+ print(tuple(out.shape)) # (10, 32)
42
+ ```
43
+ """
44
+
45
+ def __init__(
46
+ self,
47
+ in_channels: int,
48
+ out_channels: int,
49
+ num_node_types: int,
50
+ num_edge_types: int,
51
+ edge_type_emb_dim: int,
52
+ edge_dim: int,
53
+ edge_attr_emb_dim: int,
54
+ heads: int = 1,
55
+ concat: bool = True,
56
+ negative_slope: float = 0.2,
57
+ root_weight: bool = True,
58
+ bias: bool = True,
59
+ **kwargs,
60
+ ):
61
+ kwargs.setdefault("aggr", "add")
62
+ super().__init__(node_dim=0, **kwargs)
63
+
64
+ self.in_channels = in_channels
65
+ self.out_channels = out_channels
66
+ self.num_node_types = num_node_types
67
+ self.num_edge_types = num_edge_types
68
+ self.edge_type_emb_dim = edge_type_emb_dim
69
+ self.edge_dim = edge_dim
70
+ self.edge_attr_emb_dim = edge_attr_emb_dim
71
+ self.heads = heads
72
+ self.concat = concat
73
+ self.negative_slope = negative_slope
74
+ self.root_weight = root_weight
75
+ self.use_bias = bias
76
+
77
+ self.hetero_lin = HeteroLinear(in_channels, out_channels, num_node_types, bias=bias)
78
+ self.edge_type_emb = keras.layers.Embedding(num_edge_types, edge_type_emb_dim)
79
+ self.edge_attr_emb = keras.layers.Dense(edge_attr_emb_dim, use_bias=False)
80
+ self.att = keras.layers.Dense(heads, use_bias=False)
81
+ self.lin = keras.layers.Dense(out_channels, use_bias=bias)
82
+
83
+ def build(self, input_shape=None):
84
+ if not self.hetero_lin.built and self.in_channels > 0:
85
+ self.hetero_lin.build((None, self.in_channels))
86
+ if not self.edge_type_emb.built:
87
+ self.edge_type_emb.build((None,))
88
+ if not self.edge_attr_emb.built and self.edge_dim > 0:
89
+ self.edge_attr_emb.build((None, self.edge_dim))
90
+ att_dim = 2 * self.out_channels + self.edge_type_emb_dim + self.edge_attr_emb_dim
91
+ if not self.att.built:
92
+ self.att.build((None, att_dim))
93
+ lin_dim = self.out_channels + self.edge_attr_emb_dim
94
+ if not self.lin.built:
95
+ self.lin.build((None, lin_dim))
96
+ super().build(input_shape)
97
+
98
+ def call(self, x, edge_index, node_type, edge_type, edge_attr=None):
99
+ if not self.built:
100
+ self.build()
101
+
102
+ x = self.hetero_lin(x, node_type)
103
+ edge_type_emb = ops.leaky_relu(
104
+ self.edge_type_emb(edge_type),
105
+ negative_slope=self.negative_slope,
106
+ )
107
+
108
+ out = self.propagate(edge_index, x=x, edge_type_emb=edge_type_emb, edge_attr=edge_attr)
109
+
110
+ if self.concat:
111
+ if self.root_weight:
112
+ out = out + ops.expand_dims(x, axis=1)
113
+ out = ops.reshape(out, (-1, self.heads * self.out_channels))
114
+ else:
115
+ out = ops.mean(out, axis=1)
116
+ if self.root_weight:
117
+ out = out + x
118
+
119
+ return out
120
+
121
+ def message(self, x_i, x_j, edge_type_emb, edge_attr, index=None):
122
+ edge_attr = ops.leaky_relu(self.edge_attr_emb(edge_attr), negative_slope=self.negative_slope)
123
+ alpha = ops.concatenate([x_i, x_j, edge_type_emb, edge_attr], axis=-1)
124
+ alpha = ops.leaky_relu(self.att(alpha), negative_slope=self.negative_slope)
125
+ alpha = softmax(alpha, index=index)
126
+
127
+ feat = self.lin(ops.concatenate([x_j, edge_attr], axis=-1))
128
+ # feat: (E, out_channels), alpha: (E, heads)
129
+ out = ops.expand_dims(feat, axis=1) * ops.expand_dims(alpha, axis=-1)
130
+ return out
131
+
@@ -0,0 +1,128 @@
1
+ from typing import Dict, List, Optional, Tuple
2
+ import keras
3
+ from keras import ops
4
+
5
+
6
+ def group(xs: List[any], aggr: Optional[str]) -> Optional[any]:
7
+ if len(xs) == 0:
8
+ return None
9
+ elif aggr is None:
10
+ return ops.stack(xs, axis=1)
11
+ elif len(xs) == 1:
12
+ return xs[0]
13
+ elif aggr == "cat":
14
+ return ops.concatenate(xs, axis=-1)
15
+ elif aggr == "sum":
16
+ return ops.sum(ops.stack(xs, axis=0), axis=0)
17
+ elif aggr == "mean":
18
+ return ops.mean(ops.stack(xs, axis=0), axis=0)
19
+ elif aggr == "min":
20
+ return ops.min(ops.stack(xs, axis=0), axis=0)
21
+ elif aggr == "max":
22
+ return ops.max(ops.stack(xs, axis=0), axis=0)
23
+ else:
24
+ raise ValueError(f"Unsupported aggregation: {aggr}")
25
+
26
+
27
+ class HeteroConv(keras.layers.Layer):
28
+ r"""A generic wrapper for computing graph convolution on heterogeneous graphs.
29
+
30
+ Args:
31
+ convs (Dict[Tuple[str, str, str], keras.layers.Layer]): A dictionary holding a
32
+ bipartite GNN layer for each individual edge type.
33
+ aggr (str, optional): The aggregation scheme to use for grouping node
34
+ embeddings generated by different relations (:obj:`"sum"`, :obj:`"mean"`,
35
+ :obj:`"min"`, :obj:`"max"`, :obj:`"cat"`, :obj:`None`). (default: :obj:`"sum"`)
36
+
37
+ Example:
38
+ ```python
39
+ import numpy as np
40
+ from k3_node.layers import HeteroConv, SAGEConv
41
+
42
+ x_dict = {
43
+ "author": np.random.rand(3, 8).astype("float32"),
44
+ "paper": np.random.rand(4, 8).astype("float32"),
45
+ }
46
+ edge_index_dict = {
47
+ ("author", "writes", "paper"): np.array([[0, 1, 2], [1, 2, 3]]),
48
+ ("paper", "written_by", "author"): np.array([[1, 2, 3], [0, 1, 2]]),
49
+ }
50
+
51
+ # One convolution per edge type; results for the same node type are summed
52
+ layer = HeteroConv({
53
+ ("author", "writes", "paper"): SAGEConv((8, 8), 16),
54
+ ("paper", "written_by", "author"): SAGEConv((8, 8), 16),
55
+ }, aggr="sum")
56
+ out_dict = layer(x_dict, edge_index_dict)
57
+ print(tuple(out_dict["author"].shape), tuple(out_dict["paper"].shape)) # (3, 16) (4, 16)
58
+ ```
59
+ """
60
+
61
+ def __init__(
62
+ self,
63
+ convs: Dict[Tuple[str, str, str], keras.layers.Layer],
64
+ aggr: Optional[str] = "sum",
65
+ **kwargs,
66
+ ):
67
+ super().__init__(**kwargs)
68
+ self.convs = convs
69
+ self.aggr = aggr
70
+
71
+ def build(self, input_shape=None):
72
+ self.built = True
73
+
74
+ def call(self, *args_dict, **kwargs_dict) -> Dict[str, any]:
75
+ out_dict: Dict[str, List[any]] = {}
76
+
77
+ for edge_type, conv in self.convs.items():
78
+ src, rel, dst = edge_type
79
+ has_edge_level_arg = False
80
+
81
+ args = []
82
+ for value_dict in args_dict:
83
+ if edge_type in value_dict:
84
+ has_edge_level_arg = True
85
+ args.append(value_dict[edge_type])
86
+ elif src == dst and src in value_dict:
87
+ args.append(value_dict[src])
88
+ elif src in value_dict or dst in value_dict:
89
+ args.append((
90
+ value_dict.get(src, None),
91
+ value_dict.get(dst, None),
92
+ ))
93
+
94
+ kwargs = {}
95
+ for arg, value_dict in kwargs_dict.items():
96
+ if not arg.endswith("_dict"):
97
+ raise ValueError(
98
+ f"Keyword arguments in '{self.__class__.__name__}' "
99
+ f"need to end with '_dict' (got '{arg}')"
100
+ )
101
+ clean_arg = arg[:-5]
102
+ if edge_type in value_dict:
103
+ has_edge_level_arg = True
104
+ kwargs[clean_arg] = value_dict[edge_type]
105
+ elif src == dst and src in value_dict:
106
+ kwargs[clean_arg] = value_dict[src]
107
+ elif src in value_dict or dst in value_dict:
108
+ kwargs[clean_arg] = (
109
+ value_dict.get(src, None),
110
+ value_dict.get(dst, None),
111
+ )
112
+
113
+ if not has_edge_level_arg:
114
+ continue
115
+
116
+ out = conv(*args, **kwargs)
117
+
118
+ if dst not in out_dict:
119
+ out_dict[dst] = [out]
120
+ else:
121
+ out_dict[dst].append(out)
122
+
123
+ result: Dict[str, any] = {}
124
+ for key, value in out_dict.items():
125
+ result[key] = group(value, self.aggr)
126
+
127
+ return result
128
+
@@ -0,0 +1,218 @@
1
+ import math
2
+ from typing import Dict, List, Tuple, Union
3
+ from keras import ops
4
+ from k3_node.layers.conv.message_passing import MessagePassing
5
+ from k3_node.layers.conv.utils import softmax
6
+ from k3_node.layers.dense.linear import HeteroDictLinear, HeteroLinear
7
+ from k3_node.ops.creation import repeat
8
+
9
+
10
+ class HGTConv(MessagePassing):
11
+ r"""The Heterogeneous Graph Transformer (HGT) operator from the
12
+ `"Heterogeneous Graph Transformer" <https://arxiv.org/abs/2003.01332>`_ paper.
13
+
14
+ Args:
15
+ in_channels (int or Dict[str, int]): Size of each input sample of every node type.
16
+ out_channels (int): Size of each output sample.
17
+ metadata (Tuple[List[str], List[Tuple[str, str, str]]]): Node types and edge types.
18
+ heads (int, optional): Number of multi-head-attentions. (default: :obj:`1`)
19
+
20
+ Example:
21
+ ```python
22
+ import numpy as np
23
+ from k3_node.layers import HGTConv
24
+
25
+ x_dict = {
26
+ "author": np.random.rand(3, 8).astype("float32"),
27
+ "paper": np.random.rand(4, 8).astype("float32"),
28
+ }
29
+ edge_index_dict = {
30
+ ("author", "writes", "paper"): np.array([[0, 1, 2], [1, 2, 3]]),
31
+ ("paper", "written_by", "author"): np.array([[1, 2, 3], [0, 1, 2]]),
32
+ }
33
+ metadata = (list(x_dict), list(edge_index_dict)) # (node types, edge types)
34
+
35
+ layer = HGTConv(in_channels={"author": 8, "paper": 8}, out_channels=16, metadata=metadata, heads=2)
36
+ out_dict = layer(x_dict, edge_index_dict)
37
+ print(tuple(out_dict["author"].shape), tuple(out_dict["paper"].shape)) # (3, 16) (4, 16)
38
+ ```
39
+ """
40
+
41
+ def __init__(
42
+ self,
43
+ in_channels: Union[int, Dict[str, int]],
44
+ out_channels: int,
45
+ metadata: Tuple[List[str], List[Tuple[str, str, str]]],
46
+ heads: int = 1,
47
+ **kwargs,
48
+ ):
49
+ super().__init__(aggr="add", node_dim=0, **kwargs)
50
+
51
+ if out_channels % heads != 0:
52
+ raise ValueError(f"'out_channels' ({out_channels}) must be divisible by 'heads' ({heads})")
53
+
54
+ if not isinstance(in_channels, dict):
55
+ in_channels = {node_type: in_channels for node_type in metadata[0]}
56
+
57
+ self.in_channels = in_channels
58
+ self.out_channels = out_channels
59
+ self.heads = heads
60
+ self.node_types = metadata[0]
61
+ self.edge_types = metadata[1]
62
+ self.edge_types_map = {edge_type: i for i, edge_type in enumerate(metadata[1])}
63
+ self.dst_node_types = {key[-1] for key in self.edge_types}
64
+
65
+ self.kqv_lin = HeteroDictLinear(self.in_channels, self.out_channels * 3)
66
+ self.out_lin = HeteroDictLinear(self.out_channels, self.out_channels, types=self.node_types)
67
+
68
+ dim = out_channels // heads
69
+ num_types = heads * len(self.edge_types)
70
+
71
+ self.k_rel = HeteroLinear(dim, dim, num_types, bias=False, is_sorted=True)
72
+ self.v_rel = HeteroLinear(dim, dim, num_types, bias=False, is_sorted=True)
73
+
74
+ def build(self, input_shape=None):
75
+ self.skip = {}
76
+ for nt in self.node_types:
77
+ self.skip[nt] = self.add_weight(
78
+ shape=(1,),
79
+ initializer="ones",
80
+ trainable=True,
81
+ name=f"skip_{nt}",
82
+ )
83
+ self.p_rel = {}
84
+ for et in self.edge_types:
85
+ key = "__".join(et)
86
+ self.p_rel[key] = self.add_weight(
87
+ shape=(1, self.heads),
88
+ initializer="ones",
89
+ trainable=True,
90
+ name=f"p_rel_{key}",
91
+ )
92
+ super().build(input_shape)
93
+
94
+ def _cat(self, x_dict: Dict[str, any]) -> Tuple[any, Dict[str, int]]:
95
+ cumsum = 0
96
+ outs = []
97
+ offset = {}
98
+ for key, x in x_dict.items():
99
+ outs.append(x)
100
+ offset[key] = cumsum
101
+ cumsum += ops.shape(x)[0]
102
+ return ops.concatenate(outs, axis=0), offset
103
+
104
+ def _construct_src_node_feat(
105
+ self,
106
+ k_dict: Dict[str, any],
107
+ v_dict: Dict[str, any],
108
+ edge_index_dict: Dict[Tuple[str, str, str], any],
109
+ ):
110
+ cumsum = 0
111
+ num_edge_types = len(self.edge_types)
112
+ H, D = self.heads, self.out_channels // self.heads
113
+
114
+ ks = []
115
+ vs = []
116
+ type_list = []
117
+ offset = {}
118
+
119
+ for edge_type in edge_index_dict.keys():
120
+ src = edge_type[0]
121
+ N = ops.shape(k_dict[src])[0]
122
+ offset[edge_type] = cumsum
123
+ cumsum += N
124
+
125
+ edge_type_offset = self.edge_types_map[edge_type]
126
+ type_vec = (
127
+ repeat(ops.reshape(ops.arange(H, dtype="int32"), (-1, 1)), N, axis=1)
128
+ * num_edge_types
129
+ + edge_type_offset
130
+ )
131
+
132
+ type_list.append(type_vec)
133
+ ks.append(k_dict[src])
134
+ vs.append(v_dict[src])
135
+
136
+ ks_cat = ops.reshape(ops.transpose(ops.concatenate(ks, axis=0), (1, 0, 2)), (-1, D))
137
+ vs_cat = ops.reshape(ops.transpose(ops.concatenate(vs, axis=0), (1, 0, 2)), (-1, D))
138
+ type_vec_cat = ops.reshape(ops.concatenate(type_list, axis=1), (-1,))
139
+
140
+ k = self.k_rel(ks_cat, type_vec_cat)
141
+ k = ops.transpose(ops.reshape(k, (H, -1, D)), (1, 0, 2))
142
+
143
+ v = self.v_rel(vs_cat, type_vec_cat)
144
+ v = ops.transpose(ops.reshape(v, (H, -1, D)), (1, 0, 2))
145
+
146
+ return k, v, offset
147
+
148
+ def call(
149
+ self,
150
+ x_dict: Dict[str, any],
151
+ edge_index_dict: Dict[Tuple[str, str, str], any],
152
+ ) -> Dict[str, any]:
153
+ if not self.built:
154
+ self.build()
155
+
156
+ F = self.out_channels
157
+ H = self.heads
158
+ D = F // H
159
+
160
+ k_dict, q_dict, v_dict, out_dict = {}, {}, {}, {}
161
+
162
+ kqv_dict = self.kqv_lin(x_dict)
163
+ for key, val in kqv_dict.items():
164
+ k, q, v = ops.split(val, 3, axis=1)
165
+ k_dict[key] = ops.reshape(k, (-1, H, D))
166
+ q_dict[key] = ops.reshape(q, (-1, H, D))
167
+ v_dict[key] = ops.reshape(v, (-1, H, D))
168
+
169
+ q, dst_offset = self._cat(q_dict)
170
+ k, v, src_offset = self._construct_src_node_feat(k_dict, v_dict, edge_index_dict)
171
+
172
+ # Build concatenated bipartite edge index and edge attributes
173
+ edge_indices = []
174
+ edge_attrs = []
175
+ for edge_type, s_offset in src_offset.items():
176
+ e_idx = edge_index_dict[edge_type]
177
+ d_offset = dst_offset[edge_type[-1]]
178
+
179
+ row = e_idx[0] + s_offset
180
+ col = e_idx[1] + d_offset
181
+ edge_indices.append(ops.stack([row, col], axis=0))
182
+
183
+ p_val = self.p_rel["__".join(edge_type)]
184
+ num_e = ops.shape(e_idx)[1]
185
+ edge_attrs.append(repeat(p_val, num_e, axis=0))
186
+
187
+ edge_index = ops.concatenate(edge_indices, axis=1)
188
+ edge_attr = ops.concatenate(edge_attrs, axis=0)
189
+
190
+ out = self.propagate(edge_index, k=k, q=q, v=v, edge_attr=edge_attr)
191
+
192
+ for node_type, start_offset in dst_offset.items():
193
+ count = ops.shape(q_dict[node_type])[0]
194
+ if node_type in self.dst_node_types:
195
+ out_dict[node_type] = out[start_offset : start_offset + count]
196
+
197
+ transformed = {}
198
+ for k_nt, v_nt in out_dict.items():
199
+ transformed[k_nt] = ops.gelu(v_nt) if v_nt is not None else v_nt
200
+ a_dict = self.out_lin(transformed)
201
+
202
+ result = {}
203
+ for node_type, o in out_dict.items():
204
+ res = a_dict[node_type]
205
+ if ops.shape(res)[-1] == ops.shape(x_dict[node_type])[-1]:
206
+ alpha = ops.sigmoid(self.skip[node_type])
207
+ res = alpha * res + (1.0 - alpha) * x_dict[node_type]
208
+ result[node_type] = res
209
+
210
+ return result
211
+
212
+ def message(self, k_j, q_i, v_j, edge_attr, index=None):
213
+ alpha = ops.sum(q_i * k_j, axis=-1) * edge_attr
214
+ alpha = alpha / math.sqrt(ops.shape(q_i)[-1])
215
+ alpha = softmax(alpha, index=index)
216
+ out = v_j * ops.expand_dims(alpha, axis=-1)
217
+ return ops.reshape(out, (-1, self.out_channels))
218
+