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,200 @@
1
+ from typing import Dict, List, Optional
2
+
3
+ import keras
4
+ from keras import ops
5
+
6
+
7
+ class JumpingKnowledge(keras.layers.Layer):
8
+ r"""The Jumping Knowledge layer aggregation module from the
9
+ `"Representation Learning on Graphs with Jumping Knowledge Networks"
10
+ <https://arxiv.org/abs/1806.03536>`_ paper.
11
+
12
+ Jumping knowledge is performed based on either **concatenation**
13
+ (:obj:`"cat"`)
14
+
15
+ .. math::
16
+
17
+ \mathbf{x}_v^{(1)} \, \Vert \, \ldots \, \Vert \, \mathbf{x}_v^{(T)},
18
+
19
+ **max pooling** (:obj:`"max"`)
20
+
21
+ .. math::
22
+
23
+ \max \left( \mathbf{x}_v^{(1)}, \ldots, \mathbf{x}_v^{(T)} \right),
24
+
25
+ or **weighted summation**
26
+
27
+ .. math::
28
+
29
+ \sum_{t=1}^T \alpha_v^{(t)} \mathbf{x}_v^{(t)}
30
+
31
+ with attention scores :math:`\alpha_v^{(t)}` obtained from a bi-directional
32
+ LSTM (:obj:`"lstm"`).
33
+
34
+ Args:
35
+ mode (str): The aggregation scheme to use
36
+ (:obj:`"cat"`, :obj:`"max"` or :obj:`"lstm"`).
37
+ channels (int, optional): The number of channels per representation.
38
+ Needs to be only set for LSTM-style aggregation.
39
+ (default: :obj:`None`)
40
+ num_layers (int, optional): The number of layers to aggregate. Needs to
41
+ be only set for LSTM-style aggregation. (default: :obj:`None`)
42
+
43
+ Example:
44
+ ```python
45
+ import numpy as np
46
+ from k3_node.models import JumpingKnowledge
47
+
48
+ # Node representations from 3 GNN layers
49
+ xs = [np.random.rand(10, 16).astype("float32") for _ in range(3)]
50
+ print(tuple(JumpingKnowledge("cat")(xs).shape)) # (10, 48): concatenate all layers
51
+ print(tuple(JumpingKnowledge("max")(xs).shape)) # (10, 16): element-wise max
52
+ print(tuple(JumpingKnowledge("lstm", channels=16, num_layers=3)(xs).shape)) # (10, 16): attention over layers
53
+ ```
54
+ """
55
+
56
+ def __init__(
57
+ self,
58
+ mode: str,
59
+ channels: Optional[int] = None,
60
+ num_layers: Optional[int] = None,
61
+ **kwargs,
62
+ ) -> None:
63
+ super().__init__(**kwargs)
64
+ self.mode = mode.lower()
65
+ assert self.mode in ['cat', 'max', 'lstm'], \
66
+ f"mode must be 'cat', 'max', or 'lstm', got '{mode}'"
67
+
68
+ self.channels = channels
69
+ self.num_layers = num_layers
70
+
71
+ if self.mode == 'lstm':
72
+ assert channels is not None, 'channels cannot be None for lstm'
73
+ assert num_layers is not None, 'num_layers cannot be None for lstm'
74
+ lstm_units = (num_layers * channels) // 2
75
+ self.lstm = keras.layers.Bidirectional(
76
+ keras.layers.LSTM(lstm_units, return_sequences=True),
77
+ merge_mode='concat',
78
+ )
79
+ self.att = keras.layers.Dense(1, use_bias=False)
80
+ else:
81
+ self.lstm = None
82
+ self.att = None
83
+
84
+ def reset_parameters(self) -> None:
85
+ r"""Resets all learnable parameters of the module."""
86
+ # Keras layers reinitialize on next call; explicit rebuild if needed
87
+ if self.lstm is not None and self.lstm.built:
88
+ for layer in self.lstm.layers:
89
+ if hasattr(layer, 'kernel') and layer.kernel is not None:
90
+ layer.kernel.assign(
91
+ keras.initializers.GlorotUniform()(layer.kernel.shape)
92
+ )
93
+ if hasattr(layer, 'recurrent_kernel') and layer.recurrent_kernel is not None:
94
+ layer.recurrent_kernel.assign(
95
+ keras.initializers.Orthogonal()(layer.recurrent_kernel.shape)
96
+ )
97
+ if hasattr(layer, 'bias') and layer.bias is not None:
98
+ layer.bias.assign(ops.zeros(layer.bias.shape))
99
+ if self.att is not None and self.att.built:
100
+ if self.att.kernel is not None:
101
+ self.att.kernel.assign(
102
+ keras.initializers.GlorotUniform()(self.att.kernel.shape)
103
+ )
104
+
105
+ def call(self, xs: List) -> object:
106
+ r"""Forward pass.
107
+
108
+ Args:
109
+ xs (List[Tensor]): List containing the layer-wise representations.
110
+ """
111
+ if self.mode == 'cat':
112
+ return ops.concatenate(xs, axis=-1)
113
+ elif self.mode == 'max':
114
+ return ops.max(ops.stack(xs, axis=-1), axis=-1)
115
+ else: # lstm
116
+ assert self.lstm is not None and self.att is not None
117
+ x = ops.stack(xs, axis=1) # [num_nodes, num_layers, num_channels]
118
+ alpha = self.lstm(x) # [num_nodes, num_layers, 2*lstm_units]
119
+ alpha = self.att(alpha) # [num_nodes, num_layers, 1]
120
+ alpha = ops.squeeze(alpha, axis=-1) # [num_nodes, num_layers]
121
+ alpha = ops.softmax(alpha, axis=-1)
122
+ return ops.sum(x * ops.expand_dims(alpha, axis=-1), axis=1)
123
+
124
+ def __repr__(self) -> str:
125
+ if self.mode == 'lstm':
126
+ return (f'{self.__class__.__name__}({self.mode}, '
127
+ f'channels={self.channels}, layers={self.num_layers})')
128
+ return f'{self.__class__.__name__}({self.mode})'
129
+
130
+
131
+ class HeteroJumpingKnowledge(keras.layers.Layer):
132
+ r"""A heterogeneous version of the :class:`JumpingKnowledge` module.
133
+
134
+ Args:
135
+ types (List[str]): The keys of the input dictionary.
136
+ mode (str): The aggregation scheme to use
137
+ (:obj:`"cat"`, :obj:`"max"` or :obj:`"lstm"`).
138
+ channels (int, optional): The number of channels per representation.
139
+ Needs to be only set for LSTM-style aggregation.
140
+ (default: :obj:`None`)
141
+ num_layers (int, optional): The number of layers to aggregate. Needs to
142
+ be only set for LSTM-style aggregation. (default: :obj:`None`)
143
+
144
+ Example:
145
+ ```python
146
+ import numpy as np
147
+ from k3_node.models import HeteroJumpingKnowledge
148
+
149
+ # Per node type: representations from 3 GNN layers
150
+ xs_dict = {
151
+ "author": [np.random.rand(3, 16).astype("float32") for _ in range(3)],
152
+ "paper": [np.random.rand(4, 16).astype("float32") for _ in range(3)],
153
+ }
154
+ model = HeteroJumpingKnowledge(["author", "paper"], mode="cat")
155
+ out_dict = model(xs_dict)
156
+ print(tuple(out_dict["author"].shape), tuple(out_dict["paper"].shape)) # (3, 48) (4, 48)
157
+ ```
158
+ """
159
+
160
+ def __init__(
161
+ self,
162
+ types: List[str],
163
+ mode: str,
164
+ channels: Optional[int] = None,
165
+ num_layers: Optional[int] = None,
166
+ **kwargs,
167
+ ) -> None:
168
+ super().__init__(**kwargs)
169
+ self.mode = mode.lower()
170
+ self.types = list(types)
171
+
172
+ self.jk_dict = {
173
+ key: JumpingKnowledge(mode, channels, num_layers)
174
+ for key in types
175
+ }
176
+
177
+ def reset_parameters(self) -> None:
178
+ r"""Resets all learnable parameters of the module."""
179
+ for jk in self.jk_dict.values():
180
+ jk.reset_parameters()
181
+
182
+ def call(self, xs_dict: Dict[str, List]) -> Dict[str, object]:
183
+ r"""Forward pass.
184
+
185
+ Args:
186
+ xs_dict (Dict[str, List[Tensor]]): A dictionary holding a
187
+ list of layer-wise representation for each type.
188
+ """
189
+ return {key: self.jk_dict[key](xs_dict[key]) for key in self.types}
190
+
191
+ def __repr__(self) -> str:
192
+ if self.mode == 'lstm':
193
+ jk = next(iter(self.jk_dict.values()))
194
+ return (f'{self.__class__.__name__}('
195
+ f'num_types={len(self.jk_dict)}, '
196
+ f'mode={self.mode}, channels={jk.channels}, '
197
+ f'layers={jk.num_layers})')
198
+ return (f'{self.__class__.__name__}(num_types={len(self.jk_dict)}, '
199
+ f'mode={self.mode})')
200
+
@@ -0,0 +1,110 @@
1
+ from typing import Callable, Optional
2
+
3
+ import keras
4
+ from keras import ops
5
+
6
+ from k3_node.layers.conv.message_passing import MessagePassing
7
+ from k3_node.layers.conv.utils import gcn_norm
8
+
9
+
10
+ class LabelPropagation(MessagePassing):
11
+ r"""The label propagation operator, firstly introduced in the
12
+ `"Learning from Labeled and Unlabeled Data with Label Propagation"
13
+ <http://mlg.eng.cam.ac.uk/zoubin/papers/CMU-CALD-02-107.pdf>`_ paper.
14
+
15
+ .. math::
16
+ \mathbf{Y}^{\prime} = \alpha \cdot \mathbf{D}^{-1/2} \mathbf{A}
17
+ \mathbf{D}^{-1/2} \mathbf{Y} + (1 - \alpha) \mathbf{Y},
18
+
19
+ where unlabeled data is inferred by labeled data via propagation.
20
+ This concrete implementation here is derived from the `"Combining Label
21
+ Propagation And Simple Models Out-performs Graph Neural Networks"
22
+ <https://arxiv.org/abs/2010.13993>`_ paper.
23
+
24
+ Args:
25
+ num_layers (int): The number of propagations.
26
+ alpha (float): The :math:`\alpha` coefficient.
27
+
28
+ Example:
29
+ ```python
30
+ import numpy as np
31
+ from k3_node.models import LabelPropagation
32
+
33
+ y = np.array([0, 1, 2, 0, 1, 2]) # node labels
34
+ train_mask = np.array([True, True, True, False, False, False]) # labels known for 3 nodes
35
+ edge_index = np.array([[0, 1, 2, 3, 4, 5], [3, 4, 5, 0, 1, 2]])
36
+ model = LabelPropagation(num_layers=3, alpha=0.9)
37
+ out = model(y, edge_index, train_mask) # soft labels for every node
38
+ print(tuple(out.shape)) # (6, 3)
39
+ ```
40
+ """
41
+ def __init__(self, num_layers: int, alpha: float, **kwargs):
42
+ super().__init__(aggr='sum', **kwargs)
43
+ self.num_layers = num_layers
44
+ self.alpha = alpha
45
+
46
+ def build(self, input_shape=None):
47
+ self.built = True
48
+
49
+ # `mask` selects the labeled nodes; it is not a Keras sequence mask, so the output gets none
50
+ # (this also stops Keras from warning that the layer drops the mask).
51
+ supports_masking = True
52
+
53
+ def compute_mask(self, *args, **kwargs):
54
+ return None
55
+
56
+ def call(
57
+ self,
58
+ y,
59
+ edge_index,
60
+ mask=None,
61
+ edge_weight=None,
62
+ post_step: Optional[Callable] = None,
63
+ ):
64
+ shape = ops.shape(y)
65
+ if len(shape) == 1:
66
+ num_classes = int(ops.max(y)) + 1
67
+ y = ops.one_hot(y, num_classes)
68
+
69
+ y = ops.cast(y, dtype="float32")
70
+ out = y
71
+ if mask is not None:
72
+ mask_shape = ops.shape(mask)
73
+ # standardize_dtype: torch.bool does not compare equal to the string "bool"
74
+ if len(mask_shape) == 1 and keras.backend.standardize_dtype(mask.dtype) == "bool":
75
+ mask_expanded = ops.expand_dims(mask, axis=-1)
76
+ out = ops.where(mask_expanded, y, ops.zeros_like(y))
77
+ else:
78
+ out_zeros = ops.zeros_like(y)
79
+ out = ops.scatter_update(out_zeros, ops.expand_dims(mask, -1), ops.take(y, mask, axis=0))
80
+
81
+ if edge_weight is None:
82
+ num_nodes = ops.shape(y)[0]
83
+ edge_index, edge_weight = gcn_norm(
84
+ edge_index,
85
+ edge_weight=None,
86
+ num_nodes=num_nodes,
87
+ add_self_loops=False,
88
+ dtype=y.dtype,
89
+ )
90
+
91
+ res = (1.0 - self.alpha) * out
92
+ for _ in range(self.num_layers):
93
+ out = self.propagate(edge_index, x=out, edge_weight=edge_weight)
94
+ out = self.alpha * out + res
95
+ if post_step is not None:
96
+ out = post_step(out)
97
+ else:
98
+ out = ops.clip(out, 0.0, 1.0)
99
+
100
+ return out
101
+
102
+ def message(self, x_j, edge_weight=None):
103
+ if edge_weight is None:
104
+ return x_j
105
+ return ops.expand_dims(edge_weight, axis=-1) * x_j
106
+
107
+ def __repr__(self) -> str:
108
+ return (f'{self.__class__.__name__}(num_layers={self.num_layers}, '
109
+ f'alpha={self.alpha})')
110
+
@@ -0,0 +1,171 @@
1
+ from typing import Optional, Union
2
+
3
+ import keras
4
+ from keras import ops
5
+ import numpy as np
6
+
7
+ from k3_node.layers.conv import LGConv
8
+
9
+
10
+ class BPRLoss:
11
+ r"""The Bayesian Personalized Ranking (BPR) loss."""
12
+ def __init__(self, lambda_reg: float = 0.0):
13
+ self.lambda_reg = lambda_reg
14
+
15
+ def __call__(self, positives, negatives, parameters=None):
16
+ diff = positives - negatives
17
+ # log(sigmoid(x)) = -log(1 + exp(-x)) or ops.log_sigmoid if available
18
+ log_prob = ops.mean(-ops.softplus(-diff))
19
+
20
+ regularization = 0.0
21
+ if self.lambda_reg != 0.0 and parameters is not None:
22
+ regularization = self.lambda_reg * ops.sum(ops.square(parameters))
23
+ regularization = regularization / ops.cast(ops.shape(positives)[0], "float32")
24
+
25
+ return -log_prob + regularization
26
+
27
+
28
+ class LightGCN(keras.Model):
29
+ r"""The LightGCN model from the `"LightGCN: Simplifying and Powering
30
+ Graph Convolution Network for Recommendation"
31
+ <https://arxiv.org/abs/2002.02126>`_ paper.
32
+
33
+ Args:
34
+ num_nodes (int): The number of nodes in the graph.
35
+ embedding_dim (int): The dimensionality of node embeddings.
36
+ num_layers (int): The number of :class:`LGConv` layers.
37
+ alpha (float or Tensor, optional): The scalar or vector specifying
38
+ the re-weighting coefficients for aggregating the final embedding.
39
+ (default: :obj:`None`)
40
+
41
+ Example:
42
+ ```python
43
+ import numpy as np
44
+ from k3_node.models import LightGCN
45
+
46
+ edge_index = np.array([[0, 1, 2, 3, 4, 5, 6, 7], [1, 2, 3, 4, 5, 6, 7, 0]]) # user-item interactions
47
+ edge_label_index = np.array([[0, 1, 2, 3], [4, 5, 6, 7]]) # pairs to score
48
+ model = LightGCN(num_nodes=50, embedding_dim=16, num_layers=2)
49
+ scores = model(edge_index, edge_label_index)
50
+ print(tuple(scores.shape)) # (4,)
51
+ print(tuple(model.recommend(edge_index, k=2).shape)) # (50, 2): top-2 recommendations per node
52
+ ```
53
+ """
54
+ def __init__(
55
+ self,
56
+ num_nodes: int,
57
+ embedding_dim: int,
58
+ num_layers: int,
59
+ alpha: Optional[Union[float, list]] = None,
60
+ **kwargs,
61
+ ):
62
+ super().__init__()
63
+
64
+ self.num_nodes = num_nodes
65
+ self.embedding_dim = embedding_dim
66
+ self.num_layers = num_layers
67
+
68
+ if alpha is None:
69
+ alpha = [1.0 / (num_layers + 1)] * (num_layers + 1)
70
+ elif isinstance(alpha, (int, float)):
71
+ alpha = [float(alpha)] * (num_layers + 1)
72
+ self.alpha_list = list(alpha)
73
+
74
+ self.embedding = keras.layers.Embedding(
75
+ num_nodes,
76
+ embedding_dim,
77
+ embeddings_initializer=keras.initializers.GlorotUniform(),
78
+ )
79
+ self.convs = [LGConv(**kwargs) for _ in range(num_layers)]
80
+
81
+ def build(self, input_shape=None):
82
+ self.embedding.build((None,))
83
+ self.built = True
84
+
85
+ def reset_parameters(self):
86
+ r"""Resets all learnable parameters of the module."""
87
+ if self.embedding.built:
88
+ self.embedding.embeddings.assign(
89
+ keras.initializers.GlorotUniform()(shape=(self.num_nodes, self.embedding_dim))
90
+ )
91
+ for conv in self.convs:
92
+ if hasattr(conv, "reset_parameters"):
93
+ conv.reset_parameters()
94
+
95
+ def get_embedding(self, edge_index, edge_weight=None):
96
+ r"""Returns the embedding of nodes in the graph."""
97
+ if not self.embedding.built:
98
+ self.embedding.build((None,))
99
+ # Embedding weights: shape [num_nodes, embedding_dim]
100
+ x = self.embedding.weights[0]
101
+ out = x * self.alpha_list[0]
102
+
103
+ for i in range(self.num_layers):
104
+ x = self.convs[i](x, edge_index, edge_weight=edge_weight)
105
+ out = out + x * self.alpha_list[i + 1]
106
+
107
+ return out
108
+
109
+ def call(self, edge_index, edge_label_index=None, edge_weight=None):
110
+ r"""Computes rankings for pairs of nodes."""
111
+ if edge_label_index is None:
112
+ edge_label_index = edge_index
113
+
114
+ out = self.get_embedding(edge_index, edge_weight)
115
+
116
+ out_src = ops.take(out, edge_label_index[0], axis=0)
117
+ out_dst = ops.take(out, edge_label_index[1], axis=0)
118
+
119
+ return ops.sum(out_src * out_dst, axis=-1)
120
+
121
+ def predict_link(
122
+ self,
123
+ edge_index,
124
+ edge_label_index=None,
125
+ edge_weight=None,
126
+ prob: bool = False,
127
+ ):
128
+ pred = ops.sigmoid(self(edge_index, edge_label_index, edge_weight))
129
+ return pred if prob else ops.round(pred)
130
+
131
+ def recommend(
132
+ self,
133
+ edge_index,
134
+ edge_weight=None,
135
+ src_index=None,
136
+ dst_index=None,
137
+ k: int = 1,
138
+ sorted: bool = True,
139
+ ):
140
+ out = self.get_embedding(edge_index, edge_weight)
141
+ out_src = ops.take(out, src_index, axis=0) if src_index is not None else out
142
+ out_dst = ops.take(out, dst_index, axis=0) if dst_index is not None else out
143
+
144
+ pred = out_src @ ops.transpose(out_dst)
145
+ top_indices = ops.top_k(pred, k=k, sorted=sorted)[1]
146
+
147
+ if dst_index is not None:
148
+ top_indices = ops.take(dst_index, top_indices, axis=0)
149
+
150
+ return top_indices
151
+
152
+ def link_pred_loss(self, pred, edge_label):
153
+ loss_fn = keras.losses.BinaryCrossentropy(from_logits=True)
154
+ return loss_fn(edge_label, pred)
155
+
156
+ def recommendation_loss(
157
+ self,
158
+ pos_edge_rank,
159
+ neg_edge_rank,
160
+ node_id=None,
161
+ lambda_reg: float = 1e-4,
162
+ ):
163
+ loss_fn = BPRLoss(lambda_reg)
164
+ emb = self.embedding.weights[0]
165
+ emb = emb if node_id is None else ops.take(emb, node_id, axis=0)
166
+ return loss_fn(pos_edge_rank, neg_edge_rank, emb)
167
+
168
+ def __repr__(self) -> str:
169
+ return (f'{self.__class__.__name__}({self.num_nodes}, '
170
+ f'{self.embedding_dim}, num_layers={self.num_layers})')
171
+
@@ -0,0 +1,181 @@
1
+ import math
2
+ from typing import Optional
3
+
4
+ import keras
5
+ from keras import ops
6
+
7
+ from k3_node.layers.conv.message_passing import MessagePassing
8
+ from k3_node.layers.conv.utils import scatter
9
+ from k3_node.layers.norm import BatchNorm
10
+ from k3_node.models.mlp import MLP
11
+
12
+
13
+ class SparseLinear(keras.layers.Layer):
14
+ r"""A sparse linear transformation operator computing :math:`\mathbf{A}\mathbf{W} + \mathbf{b}`.
15
+
16
+ Example:
17
+ ```python
18
+ import numpy as np
19
+ from k3_node.models import SparseLinear
20
+
21
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
22
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
23
+
24
+ layer = SparseLinear(in_channels=10, out_channels=16) # in_channels = number of nodes
25
+ out = layer(edge_index) # multiplies the adjacency matrix with a learned weight matrix
26
+ print(tuple(out.shape)) # (10, 16)
27
+ ```
28
+ """
29
+ def __init__(self, in_channels: int, out_channels: int, bias: bool = True, **kwargs):
30
+ super().__init__(**kwargs)
31
+ self.in_channels = in_channels
32
+ self.out_channels = out_channels
33
+ self.use_bias = bias
34
+
35
+ self.weight = self.add_weight(
36
+ shape=(in_channels, out_channels),
37
+ initializer=keras.initializers.GlorotUniform(),
38
+ trainable=True,
39
+ name="weight",
40
+ )
41
+ if bias:
42
+ self.bias = self.add_weight(
43
+ shape=(out_channels,),
44
+ initializer="zeros",
45
+ trainable=True,
46
+ name="bias",
47
+ )
48
+ else:
49
+ self.bias = None
50
+
51
+ def build(self, input_shape=None):
52
+ self.built = True
53
+
54
+ def reset_parameters(self):
55
+ self.weight.assign(
56
+ keras.initializers.GlorotUniform()(shape=(self.in_channels, self.out_channels))
57
+ )
58
+ if self.use_bias and self.bias is not None:
59
+ self.bias.assign(ops.zeros((self.out_channels,)))
60
+
61
+ def call(self, edge_index, edge_weight=None):
62
+ row, col = edge_index[0], edge_index[1]
63
+ weight_j = ops.take(self.weight, row, axis=0)
64
+ if edge_weight is not None:
65
+ weight_j = ops.expand_dims(edge_weight, -1) * weight_j
66
+
67
+ out = scatter(weight_j, col, dim=0, dim_size=self.in_channels, reduce="sum")
68
+ if self.use_bias and self.bias is not None:
69
+ out = out + self.bias
70
+ return out
71
+
72
+
73
+ class LINKX(keras.Model):
74
+ r"""The LINKX model from the `"Large Scale Learning on Non-Homophilous
75
+ Graphs: New Benchmarks and Strong Simple Methods"
76
+ <https://arxiv.org/abs/2110.14446>`_ paper.
77
+
78
+ Args:
79
+ num_nodes (int): The number of nodes in the graph.
80
+ in_channels (int): Size of each input sample.
81
+ hidden_channels (int): Size of each hidden sample.
82
+ out_channels (int): Size of each output sample.
83
+ num_layers (int): Number of layers of :math:`\textrm{MLP}_{f}`.
84
+ num_edge_layers (int, optional): Number of layers of
85
+ :math:`\textrm{MLP}_{\mathbf{A}}`. (default: :obj:`1`)
86
+ num_node_layers (int, optional): Number of layers of
87
+ :math:`\textrm{MLP}_{\mathbf{X}}`. (default: :obj:`1`)
88
+ dropout (float, optional): Dropout probability. (default: :obj:`0.0`)
89
+
90
+ Example:
91
+ ```python
92
+ import numpy as np
93
+ from k3_node.models import LINKX
94
+
95
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
96
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
97
+
98
+ # LINKX learns from node features and adjacency separately, which suits heterophilic graphs
99
+ model = LINKX(num_nodes=10, in_channels=8, hidden_channels=32, out_channels=4, num_layers=2)
100
+ out = model(x, edge_index)
101
+ print(tuple(out.shape)) # (10, 4)
102
+ ```
103
+ """
104
+ def __init__(
105
+ self,
106
+ num_nodes: int,
107
+ in_channels: int,
108
+ hidden_channels: int,
109
+ out_channels: int,
110
+ num_layers: int,
111
+ num_edge_layers: int = 1,
112
+ num_node_layers: int = 1,
113
+ dropout: float = 0.0,
114
+ **kwargs,
115
+ ):
116
+ super().__init__(**kwargs)
117
+
118
+ self.num_nodes = num_nodes
119
+ self.in_channels = in_channels
120
+ self.hidden_channels = hidden_channels
121
+ self.out_channels = out_channels
122
+ self.num_edge_layers = num_edge_layers
123
+
124
+ self.edge_lin = SparseLinear(num_nodes, hidden_channels)
125
+
126
+ if num_edge_layers > 1:
127
+ self.edge_norm = BatchNorm(hidden_channels)
128
+ channels = [hidden_channels] * num_edge_layers
129
+ self.edge_mlp = MLP(channels, dropout=0.0, act_first=True)
130
+ else:
131
+ self.edge_norm = None
132
+ self.edge_mlp = None
133
+
134
+ channels = [in_channels] + [hidden_channels] * num_node_layers
135
+ self.node_mlp = MLP(channels, dropout=0.0, act_first=True)
136
+
137
+ self.cat_lin1 = keras.layers.Dense(hidden_channels)
138
+ self.cat_lin2 = keras.layers.Dense(hidden_channels)
139
+
140
+ channels = [hidden_channels] * num_layers + [out_channels]
141
+ self.final_mlp = MLP(channels, dropout=dropout, act_first=True)
142
+
143
+ def build(self, input_shape=None):
144
+ self.edge_lin.build()
145
+ self.built = True
146
+
147
+ def reset_parameters(self):
148
+ r"""Resets all learnable parameters of the module."""
149
+ self.edge_lin.reset_parameters()
150
+ if self.edge_norm is not None and hasattr(self.edge_norm, "reset_parameters"):
151
+ self.edge_norm.reset_parameters()
152
+ if self.edge_mlp is not None and hasattr(self.edge_mlp, "reset_parameters"):
153
+ self.edge_mlp.reset_parameters()
154
+ self.node_mlp.reset_parameters()
155
+ self.final_mlp.reset_parameters()
156
+
157
+ def call(self, x, edge_index=None, edge_weight=None, training=None):
158
+ if edge_index is None and isinstance(x, (tuple, list)):
159
+ if len(x) >= 2:
160
+ x, edge_index = x[0], x[1]
161
+ out = self.edge_lin(edge_index, edge_weight)
162
+
163
+ if self.edge_norm is not None and self.edge_mlp is not None:
164
+ out = ops.relu(out)
165
+ out = self.edge_norm(out, training=training)
166
+ out = self.edge_mlp(out, training=training)
167
+
168
+ out = out + self.cat_lin1(out)
169
+
170
+ if x is not None:
171
+ x = self.node_mlp(x, training=training)
172
+ out = out + x
173
+ out = out + self.cat_lin2(x)
174
+
175
+ return self.final_mlp(ops.relu(out), training=training)
176
+
177
+ def __repr__(self) -> str:
178
+ return (f'{self.__class__.__name__}(num_nodes={self.num_nodes}, '
179
+ f'in_channels={self.in_channels}, '
180
+ f'out_channels={self.out_channels})')
181
+