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,954 @@
1
+ import os
2
+ from typing import Optional, Union, Tuple, List, Callable
3
+ import numpy as np
4
+
5
+ import keras
6
+ from keras import layers, ops
7
+
8
+ from k3_node.layers.conv.utils import softmax
9
+ from k3_node.data.download import download_google_url
10
+ from k3_node.ops.segment import segment_sum
11
+
12
+
13
+ def sce_loss(x, y, alpha: float = 3.0):
14
+ r"""Scaled Cosine Error (SCE) loss from `"GraphMAE: Masked Autoencoding for Graph
15
+ Self-Supervised Learning" <https://arxiv.org/abs/2205.10803>`_ and GraphMAE2.
16
+
17
+ Args:
18
+ x (Tensor): Predicted node representations.
19
+ y (Tensor): Target node representations.
20
+ alpha (float, optional): Scaling exponent. (default: ``3.0``)
21
+ """
22
+ x_norm = ops.sqrt(ops.sum(ops.power(x, 2), axis=-1, keepdims=True) + 1e-12)
23
+ x = x / x_norm
24
+ y_norm = ops.sqrt(ops.sum(ops.power(y, 2), axis=-1, keepdims=True) + 1e-12)
25
+ y = y / y_norm
26
+
27
+ cos_sim = ops.sum(x * y, axis=-1)
28
+ diff = ops.clip(1.0 - cos_sim, 0.0, 2.0)
29
+ loss = ops.power(diff, alpha)
30
+ return ops.mean(loss)
31
+
32
+
33
+ def _get_activation(name: Optional[Union[str, Callable]]):
34
+ if name is None:
35
+ return None
36
+ if isinstance(name, str):
37
+ name_lower = name.lower()
38
+ if name_lower == "prelu":
39
+ return layers.PReLU(shared_axes=[1])
40
+ elif name_lower == "relu":
41
+ return layers.ReLU()
42
+ elif name_lower == "gelu":
43
+ return layers.Activation("gelu")
44
+ elif name_lower == "silu":
45
+ return layers.Activation("silu")
46
+ elif name_lower == "elu":
47
+ return layers.ELU()
48
+ else:
49
+ return layers.Activation(name)
50
+ elif isinstance(name, layers.Layer):
51
+ return name
52
+ elif callable(name):
53
+ return layers.Activation(name)
54
+ return None
55
+
56
+
57
+ def _get_norm(name: Optional[str], dim: int):
58
+ if name is None:
59
+ return None
60
+ name_lower = name.lower()
61
+ if name_lower in ("layernorm", "layer_norm"):
62
+ return layers.LayerNormalization(axis=-1, epsilon=1e-5)
63
+ elif name_lower in ("batchnorm", "batch_norm"):
64
+ return layers.BatchNormalization(axis=-1, momentum=0.9, epsilon=1e-5)
65
+ return None
66
+
67
+
68
+ class GraphMAE2GATConv(layers.Layer):
69
+ r"""GAT convolution layer matching GraphMAE2's architecture."""
70
+
71
+ def __init__(
72
+ self,
73
+ in_feats: int,
74
+ out_feats: int,
75
+ num_heads: int,
76
+ feat_drop: float = 0.0,
77
+ attn_drop: float = 0.0,
78
+ negative_slope: float = 0.2,
79
+ residual: bool = False,
80
+ activation: Optional[Union[str, Callable]] = None,
81
+ bias: bool = True,
82
+ norm: Optional[str] = None,
83
+ concat_out: bool = True,
84
+ **kwargs,
85
+ ):
86
+ super().__init__(**kwargs)
87
+ self.in_feats = in_feats
88
+ self.out_feats = out_feats
89
+ self.num_heads = num_heads
90
+ self.feat_drop_rate = feat_drop
91
+ self.attn_drop_rate = attn_drop
92
+ self.negative_slope = negative_slope
93
+ self.use_residual = residual
94
+ self.concat_out = concat_out
95
+ self.use_bias = bias
96
+ self.norm_name = norm
97
+ self.act_name = activation
98
+
99
+ self.fc = layers.Dense(num_heads * out_feats, use_bias=False)
100
+ self.feat_drop = layers.Dropout(feat_drop) if feat_drop > 0.0 else None
101
+ self.attn_drop = layers.Dropout(attn_drop) if attn_drop > 0.0 else None
102
+
103
+ if residual and in_feats != num_heads * out_feats:
104
+ self.res_fc = layers.Dense(num_heads * out_feats, use_bias=False)
105
+ else:
106
+ self.res_fc = None
107
+
108
+ total_dim = num_heads * out_feats if concat_out else out_feats
109
+ self.norm = _get_norm(norm, total_dim)
110
+ self.activation = _get_activation(activation)
111
+
112
+ def build(self, input_shape=None):
113
+ shape = input_shape or (None, self.in_feats)
114
+ in_dim = shape[-1] if shape is not None and shape[-1] is not None else self.in_feats
115
+
116
+ self.fc.build((None, in_dim))
117
+ if self.res_fc is not None:
118
+ self.res_fc.build((None, in_dim))
119
+
120
+ self.attn_l = self.add_weight(
121
+ shape=(1, self.num_heads, self.out_feats),
122
+ initializer="glorot_uniform",
123
+ trainable=True,
124
+ name="attn_l",
125
+ )
126
+ self.attn_r = self.add_weight(
127
+ shape=(1, self.num_heads, self.out_feats),
128
+ initializer="glorot_uniform",
129
+ trainable=True,
130
+ name="attn_r",
131
+ )
132
+
133
+ if self.use_bias:
134
+ self.bias = self.add_weight(
135
+ shape=(self.num_heads * self.out_feats,),
136
+ initializer="zeros",
137
+ trainable=True,
138
+ name="bias",
139
+ )
140
+ else:
141
+ self.bias = None
142
+
143
+ total_dim = self.num_heads * self.out_feats if self.concat_out else self.out_feats
144
+ if self.norm is not None:
145
+ self.norm.build((None, total_dim))
146
+ if self.activation is not None and hasattr(self.activation, "build"):
147
+ self.activation.build((None, total_dim))
148
+
149
+ self.built = True
150
+
151
+ def call(self, x, edge_index, training=False):
152
+ h = self.feat_drop(x, training=training) if self.feat_drop is not None else x
153
+ feat_src = ops.reshape(self.fc(h), (-1, self.num_heads, self.out_feats))
154
+ feat_dst = feat_src
155
+
156
+ el = ops.sum(feat_src * self.attn_l, axis=-1, keepdims=True)
157
+ er = ops.sum(feat_dst * self.attn_r, axis=-1, keepdims=True)
158
+
159
+ row = ops.cast(edge_index[0], "int32")
160
+ col = ops.cast(edge_index[1], "int32")
161
+
162
+ el_src = ops.take(el, row, axis=0)
163
+ er_dst = ops.take(er, col, axis=0)
164
+ e = ops.leaky_relu(el_src + er_dst, negative_slope=self.negative_slope)
165
+
166
+ num_nodes = ops.shape(x)[0]
167
+ a = softmax(e, col, num_nodes=num_nodes, dim=0)
168
+ if self.attn_drop is not None:
169
+ a = self.attn_drop(a, training=training)
170
+
171
+ msg = a * ops.take(feat_src, row, axis=0)
172
+ rst = segment_sum(msg, col, num_segments=num_nodes)
173
+
174
+ if self.bias is not None:
175
+ rst = rst + ops.reshape(self.bias, (1, self.num_heads, self.out_feats))
176
+
177
+ if self.res_fc is not None:
178
+ rst = rst + ops.reshape(self.res_fc(x), (num_nodes, self.num_heads, self.out_feats))
179
+
180
+ if self.concat_out:
181
+ rst = ops.reshape(rst, (num_nodes, self.num_heads * self.out_feats))
182
+ else:
183
+ rst = ops.mean(rst, axis=1)
184
+
185
+ if self.norm is not None:
186
+ rst = self.norm(rst)
187
+
188
+ if self.activation is not None:
189
+ rst = self.activation(rst)
190
+
191
+ return rst
192
+
193
+
194
+ class GraphMAE2GAT(layers.Layer):
195
+ r"""Multi-layer GAT encoder or decoder for GraphMAE2."""
196
+
197
+ def __init__(
198
+ self,
199
+ in_dim: int,
200
+ num_hidden: int,
201
+ out_dim: int,
202
+ num_layers: int,
203
+ nhead: int,
204
+ nhead_out: int,
205
+ activation: Optional[str] = "prelu",
206
+ feat_drop: float = 0.0,
207
+ attn_drop: float = 0.0,
208
+ negative_slope: float = 0.2,
209
+ residual: bool = True,
210
+ norm: Optional[str] = "layernorm",
211
+ concat_out: bool = True,
212
+ encoding: bool = True,
213
+ **kwargs,
214
+ ):
215
+ super().__init__(**kwargs)
216
+ self.in_dim = in_dim
217
+ self.num_hidden = num_hidden
218
+ self.out_dim = out_dim
219
+ self.num_layers = num_layers
220
+ self.nhead = nhead
221
+ self.nhead_out = nhead_out
222
+ self.concat_out = concat_out
223
+ self.encoding = encoding
224
+
225
+ self.gat_layers = []
226
+
227
+ last_activation = activation if encoding else None
228
+ last_residual = (encoding and residual)
229
+ last_norm = norm if encoding else None
230
+
231
+ if num_layers == 1:
232
+ self.gat_layers.append(
233
+ GraphMAE2GATConv(
234
+ in_feats=in_dim,
235
+ out_feats=out_dim,
236
+ num_heads=nhead_out,
237
+ feat_drop=feat_drop,
238
+ attn_drop=attn_drop,
239
+ negative_slope=negative_slope,
240
+ residual=last_residual,
241
+ norm=last_norm,
242
+ activation=last_activation,
243
+ concat_out=concat_out,
244
+ )
245
+ )
246
+ else:
247
+ # Layer 0
248
+ self.gat_layers.append(
249
+ GraphMAE2GATConv(
250
+ in_feats=in_dim,
251
+ out_feats=num_hidden,
252
+ num_heads=nhead,
253
+ feat_drop=feat_drop,
254
+ attn_drop=attn_drop,
255
+ negative_slope=negative_slope,
256
+ residual=residual,
257
+ norm=norm,
258
+ activation=activation,
259
+ concat_out=concat_out,
260
+ )
261
+ )
262
+ # Intermediate layers
263
+ for _ in range(1, num_layers - 1):
264
+ self.gat_layers.append(
265
+ GraphMAE2GATConv(
266
+ in_feats=num_hidden * nhead,
267
+ out_feats=num_hidden,
268
+ num_heads=nhead,
269
+ feat_drop=feat_drop,
270
+ attn_drop=attn_drop,
271
+ negative_slope=negative_slope,
272
+ residual=residual,
273
+ norm=norm,
274
+ activation=activation,
275
+ concat_out=concat_out,
276
+ )
277
+ )
278
+ # Output layer
279
+ self.gat_layers.append(
280
+ GraphMAE2GATConv(
281
+ in_feats=num_hidden * nhead,
282
+ out_feats=out_dim,
283
+ num_heads=nhead_out,
284
+ feat_drop=feat_drop,
285
+ attn_drop=attn_drop,
286
+ negative_slope=negative_slope,
287
+ residual=last_residual,
288
+ norm=last_norm,
289
+ activation=last_activation,
290
+ concat_out=concat_out,
291
+ )
292
+ )
293
+
294
+ def build(self, input_shape=None):
295
+ for layer in self.gat_layers:
296
+ if hasattr(layer, "build") and not layer.built:
297
+ layer.build()
298
+ self.built = True
299
+
300
+ def call(self, x, edge_index, training=False):
301
+ h = x
302
+ for layer in self.gat_layers:
303
+ h = layer(h, edge_index, training=training)
304
+ return h
305
+
306
+
307
+ class GraphMAE2(layers.Layer):
308
+ r"""The GraphMAE2 model from `"GraphMAE2: A Decoding-Enhanced Masked
309
+ Self-Supervised Learning Framework for Graphs" <https://arxiv.org/abs/2304.04779>`_.
310
+
311
+ Args:
312
+ in_dim (int): Dimensionality of input node features.
313
+ num_hidden (int): Dimensionality of hidden node representations.
314
+ num_layers (int, optional): Number of encoder layers. (default: ``4``)
315
+ num_dec_layers (int, optional): Number of decoder layers. (default: ``1``)
316
+ num_remasking (int, optional): Number of remasking views in decoder. (default: ``3``)
317
+ nhead (int, optional): Number of attention heads in encoder. (default: ``8``)
318
+ nhead_out (int, optional): Number of attention heads in decoder output. (default: ``1``)
319
+ activation (str, optional): Activation function. (default: ``"prelu"``)
320
+ feat_drop (float, optional): Node feature dropout rate. (default: ``0.2``)
321
+ attn_drop (float, optional): Attention dropout rate. (default: ``0.1``)
322
+ negative_slope (float, optional): LeakyReLU negative slope. (default: ``0.2``)
323
+ residual (bool, optional): Whether to use residual connections. (default: ``True``)
324
+ norm (str, optional): Normalization type (``"layernorm"``, ``"batchnorm"``, or ``None``). (default: ``"layernorm"``)
325
+ mask_rate (float, optional): Fraction of input nodes to mask. (default: ``0.5``)
326
+ remask_rate (float, optional): Fraction of latent nodes to remask. (default: ``0.5``)
327
+ remask_method (str, optional): Remasking method (``"random"`` or ``"fixed"``). (default: ``"random"``)
328
+ loss_fn (str, optional): Reconstruction loss type (``"sce"`` or ``"mse"``). (default: ``"sce"``)
329
+ alpha_l (float, optional): Power exponent in Scaled Cosine Error loss. (default: ``2.0``)
330
+ lam (float, optional): Weight of the latent prediction loss term. (default: ``1.0``)
331
+ momentum (float, optional): Teacher EMA update momentum. (default: ``0.996``)
332
+ delayed_ema_epoch (int, optional): Epoch to begin EMA teacher updates. (default: ``0``)
333
+
334
+ Example:
335
+ ```python
336
+ import numpy as np
337
+ from k3_node.models import GraphMAE2
338
+
339
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
340
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
341
+
342
+ model = GraphMAE2(in_dim=8, num_hidden=32, num_layers=2, num_dec_layers=1, nhead=4, nhead_out=1)
343
+ print(tuple(model(x, edge_index).shape)) # (10, 32): node embeddings
344
+ loss = model.loss(x, edge_index) # masked feature reconstruction loss for pre-training
345
+ print(tuple(loss.shape)) # ()
346
+ ```
347
+ """
348
+
349
+ def __init__(
350
+ self,
351
+ in_dim: int,
352
+ num_hidden: int,
353
+ num_layers: int = 4,
354
+ num_dec_layers: int = 1,
355
+ num_remasking: int = 3,
356
+ nhead: int = 8,
357
+ nhead_out: int = 1,
358
+ activation: str = "prelu",
359
+ feat_drop: float = 0.2,
360
+ attn_drop: float = 0.1,
361
+ negative_slope: float = 0.2,
362
+ residual: bool = True,
363
+ norm: Optional[str] = "layernorm",
364
+ mask_rate: float = 0.5,
365
+ remask_rate: float = 0.5,
366
+ remask_method: str = "random",
367
+ loss_fn: str = "sce",
368
+ alpha_l: float = 2.0,
369
+ lam: float = 1.0,
370
+ momentum: float = 0.996,
371
+ delayed_ema_epoch: int = 0,
372
+ **kwargs,
373
+ ):
374
+ super().__init__(**kwargs)
375
+ self.in_dim = in_dim
376
+ self.num_hidden = num_hidden
377
+ self.num_layers = num_layers
378
+ self.num_dec_layers = num_dec_layers
379
+ self.num_remasking = num_remasking
380
+ self.nhead = nhead
381
+ self.nhead_out = nhead_out
382
+ self.mask_rate = mask_rate
383
+ self.remask_rate = remask_rate
384
+ self.remask_method = remask_method
385
+ self.loss_fn = loss_fn
386
+ self.alpha_l = alpha_l
387
+ self.lam = lam
388
+ self.momentum = momentum
389
+ self.delayed_ema_epoch = delayed_ema_epoch
390
+
391
+ assert num_hidden % nhead == 0, f"num_hidden ({num_hidden}) must be divisible by nhead ({nhead})"
392
+ assert num_hidden % nhead_out == 0, f"num_hidden ({num_hidden}) must be divisible by nhead_out ({nhead_out})"
393
+
394
+ enc_num_hidden = num_hidden // nhead
395
+ dec_in_dim = num_hidden
396
+ dec_num_hidden = num_hidden // nhead
397
+
398
+ # 1. Student Encoder
399
+ self.encoder = GraphMAE2GAT(
400
+ in_dim=in_dim,
401
+ num_hidden=enc_num_hidden,
402
+ out_dim=enc_num_hidden,
403
+ num_layers=num_layers,
404
+ nhead=nhead,
405
+ nhead_out=nhead,
406
+ activation=activation,
407
+ feat_drop=feat_drop,
408
+ attn_drop=attn_drop,
409
+ negative_slope=negative_slope,
410
+ residual=residual,
411
+ norm=norm,
412
+ concat_out=True,
413
+ encoding=True,
414
+ )
415
+
416
+ # 2. Decoder
417
+ self.decoder = GraphMAE2GAT(
418
+ in_dim=dec_in_dim,
419
+ num_hidden=dec_num_hidden,
420
+ out_dim=in_dim,
421
+ num_layers=num_dec_layers,
422
+ nhead=nhead,
423
+ nhead_out=nhead_out,
424
+ activation=activation,
425
+ feat_drop=feat_drop,
426
+ attn_drop=attn_drop,
427
+ negative_slope=negative_slope,
428
+ residual=residual,
429
+ norm=norm,
430
+ concat_out=True,
431
+ encoding=False,
432
+ )
433
+
434
+ self.encoder_to_decoder = layers.Dense(dec_in_dim, use_bias=False)
435
+
436
+ # 3. Projector & Predictor
437
+ self.projector = keras.Sequential([
438
+ layers.Dense(256),
439
+ layers.PReLU(shared_axes=[1]),
440
+ layers.Dense(num_hidden),
441
+ ])
442
+
443
+ self.predictor = keras.Sequential([
444
+ layers.PReLU(shared_axes=[1]),
445
+ layers.Dense(num_hidden),
446
+ ])
447
+
448
+ # 4. Teacher EMA networks
449
+ self.encoder_ema = GraphMAE2GAT(
450
+ in_dim=in_dim,
451
+ num_hidden=enc_num_hidden,
452
+ out_dim=enc_num_hidden,
453
+ num_layers=num_layers,
454
+ nhead=nhead,
455
+ nhead_out=nhead,
456
+ activation=activation,
457
+ feat_drop=feat_drop,
458
+ attn_drop=attn_drop,
459
+ negative_slope=negative_slope,
460
+ residual=residual,
461
+ norm=norm,
462
+ concat_out=True,
463
+ encoding=True,
464
+ trainable=False,
465
+ )
466
+
467
+ self.projector_ema = keras.Sequential([
468
+ layers.Dense(256, trainable=False),
469
+ layers.PReLU(shared_axes=[1], trainable=False),
470
+ layers.Dense(num_hidden, trainable=False),
471
+ ], trainable=False)
472
+
473
+ self.enc_mask_token = self.add_weight(
474
+ shape=(1, self.in_dim),
475
+ initializer="glorot_normal",
476
+ trainable=True,
477
+ name="enc_mask_token",
478
+ )
479
+ self.dec_mask_token = self.add_weight(
480
+ shape=(1, self.num_hidden),
481
+ initializer="glorot_normal",
482
+ trainable=True,
483
+ name="dec_mask_token",
484
+ )
485
+
486
+ def build(self, input_shape=None):
487
+ if not hasattr(self, "enc_mask_token") or self.enc_mask_token is None:
488
+ self.enc_mask_token = self.add_weight(
489
+ shape=(1, self.in_dim),
490
+ initializer="glorot_normal",
491
+ trainable=True,
492
+ name="enc_mask_token",
493
+ )
494
+ if not hasattr(self, "dec_mask_token") or self.dec_mask_token is None:
495
+ self.dec_mask_token = self.add_weight(
496
+ shape=(1, self.num_hidden),
497
+ initializer="glorot_normal",
498
+ trainable=True,
499
+ name="dec_mask_token",
500
+ )
501
+
502
+ self.encoder.build((None, self.in_dim))
503
+ self.encoder_to_decoder.build((None, self.num_hidden))
504
+ self.decoder.build((None, self.num_hidden))
505
+ self.projector.build((None, self.num_hidden))
506
+ self.predictor.build((None, self.num_hidden))
507
+ self.encoder_ema.build((None, self.in_dim))
508
+ self.projector_ema.build((None, self.num_hidden))
509
+
510
+ # Copy initial weights from student to teacher
511
+ for p_s, p_t in zip(self.encoder.weights, self.encoder_ema.weights):
512
+ p_t.assign(p_s)
513
+ for p_s, p_t in zip(self.projector.weights, self.projector_ema.weights):
514
+ p_t.assign(p_s)
515
+
516
+ self.built = True
517
+
518
+ def embed(self, x, edge_index, training=None):
519
+ r"""Generates node embeddings with the encoder."""
520
+ if not self.built:
521
+ self.build((None, self.in_dim))
522
+ return self.encoder(x, edge_index, training=training)
523
+
524
+ def encoding_mask_noise(self, x, mask_rate: Optional[float] = None, mask_nodes=None):
525
+ r"""Masks node features for encoder input."""
526
+ x = ops.convert_to_tensor(x) # NumPy inputs cannot be mixed with backend tensors
527
+ rate = self.mask_rate if mask_rate is None else mask_rate
528
+ num_nodes = ops.shape(x)[0]
529
+
530
+ if mask_nodes is None:
531
+ perm = np.random.permutation(num_nodes)
532
+ num_mask_nodes = int(rate * num_nodes)
533
+ mask_nodes = ops.convert_to_tensor(perm[:num_mask_nodes], dtype="int32")
534
+ keep_nodes = ops.convert_to_tensor(perm[num_mask_nodes:], dtype="int32")
535
+ else:
536
+ mask_nodes = ops.cast(mask_nodes, "int32")
537
+ all_mask = np.zeros(num_nodes, dtype=bool)
538
+ all_mask[ops.convert_to_numpy(mask_nodes)] = True
539
+ keep_nodes = ops.convert_to_tensor(np.where(~all_mask)[0], dtype="int32")
540
+
541
+ # Replace masked nodes with enc_mask_token
542
+ # Create a zeroed masked version
543
+ mask_vector = np.zeros(num_nodes, dtype=np.float32)
544
+ mask_vector[ops.convert_to_numpy(mask_nodes)] = 1.0
545
+ mask_tensor = ops.expand_dims(ops.convert_to_tensor(mask_vector, dtype=x.dtype), -1)
546
+
547
+ masked_x = x * (1.0 - mask_tensor) + mask_tensor * self.enc_mask_token
548
+ return masked_x, mask_nodes, keep_nodes
549
+
550
+ def random_remask(self, rep, remask_rate: Optional[float] = None, remask_nodes=None):
551
+ r"""Remasks latent representation for decoder input."""
552
+ rate = self.remask_rate if remask_rate is None else remask_rate
553
+ num_nodes = ops.shape(rep)[0]
554
+
555
+ if remask_nodes is None:
556
+ perm = np.random.permutation(num_nodes)
557
+ num_remask_nodes = int(rate * num_nodes)
558
+ remask_nodes = ops.convert_to_tensor(perm[:num_remask_nodes], dtype="int32")
559
+ rekeep_nodes = ops.convert_to_tensor(perm[num_remask_nodes:], dtype="int32")
560
+ else:
561
+ remask_nodes = ops.cast(remask_nodes, "int32")
562
+ all_mask = np.zeros(num_nodes, dtype=bool)
563
+ all_mask[ops.convert_to_numpy(remask_nodes)] = True
564
+ rekeep_nodes = ops.convert_to_tensor(np.where(~all_mask)[0], dtype="int32")
565
+
566
+ remask_vector = np.zeros(num_nodes, dtype=np.float32)
567
+ remask_vector[ops.convert_to_numpy(remask_nodes)] = 1.0
568
+ remask_tensor = ops.expand_dims(ops.convert_to_tensor(remask_vector, dtype=rep.dtype), -1)
569
+
570
+ remasked_rep = rep * (1.0 - remask_tensor) + remask_tensor * self.dec_mask_token
571
+ return remasked_rep, remask_nodes, rekeep_nodes
572
+
573
+ def ema_update(self, momentum: Optional[float] = None):
574
+ r"""Updates teacher EMA parameters."""
575
+ m = self.momentum if momentum is None else momentum
576
+ for p_s, p_t in zip(self.encoder.weights, self.encoder_ema.weights):
577
+ p_t.assign(p_t * m + p_s * (1.0 - m))
578
+ for p_s, p_t in zip(self.projector.weights, self.projector_ema.weights):
579
+ p_t.assign(p_t * m + p_s * (1.0 - m))
580
+
581
+ def loss(
582
+ self,
583
+ x,
584
+ edge_index,
585
+ mask_nodes=None,
586
+ targets=None,
587
+ epoch: int = 0,
588
+ training: bool = True,
589
+ ):
590
+ r"""Computes GraphMAE2 loss: attribute reconstruction loss + latent prediction loss."""
591
+ if not self.built:
592
+ self.build((None, self.in_dim))
593
+
594
+ # 1. Masking
595
+ masked_x, mask_nodes, keep_nodes = self.encoding_mask_noise(x, mask_nodes=mask_nodes)
596
+
597
+ # 2. Student encoder
598
+ enc_rep = self.encoder(masked_x, edge_index, training=training)
599
+
600
+ # 3. Teacher EMA target (no gradient)
601
+ teacher_rep = ops.stop_gradient(self.encoder_ema(x, edge_index, training=False))
602
+ if targets is not None:
603
+ latent_target = ops.stop_gradient(self.projector_ema(ops.take(teacher_rep, targets, axis=0)))
604
+ latent_pred = self.predictor(self.projector(ops.take(enc_rep, targets, axis=0)))
605
+ else:
606
+ latent_target = ops.stop_gradient(self.projector_ema(ops.take(teacher_rep, keep_nodes, axis=0)))
607
+ latent_pred = self.predictor(self.projector(ops.take(enc_rep, keep_nodes, axis=0)))
608
+
609
+ loss_latent = sce_loss(latent_pred, latent_target, alpha=1.0)
610
+
611
+ # 4. Decoder attribute reconstruction
612
+ origin_rep = self.encoder_to_decoder(enc_rep)
613
+
614
+ criterion = sce_loss if self.loss_fn == "sce" else (lambda pred, tgt: ops.mean(ops.power(pred - tgt, 2)))
615
+
616
+ loss_rec_all = 0.0
617
+ if self.remask_method == "random":
618
+ for _ in range(self.num_remasking):
619
+ rep, _, _ = self.random_remask(origin_rep)
620
+ recon = self.decoder(rep, edge_index, training=training)
621
+ x_init = ops.take(x, mask_nodes, axis=0)
622
+ x_rec = ops.take(recon, mask_nodes, axis=0)
623
+ loss_rec_all = loss_rec_all + criterion(x_rec, x_init, alpha=self.alpha_l) if self.loss_fn == "sce" else loss_rec_all + criterion(x_rec, x_init)
624
+ loss_rec = loss_rec_all / float(self.num_remasking)
625
+ else:
626
+ # Fixed remasking
627
+ mask_vector = np.zeros(ops.shape(x)[0], dtype=np.float32)
628
+ mask_vector[ops.convert_to_numpy(mask_nodes)] = 1.0
629
+ mask_tensor = ops.expand_dims(ops.convert_to_tensor(mask_vector, dtype=origin_rep.dtype), -1)
630
+ rep = origin_rep * (1.0 - mask_tensor)
631
+ recon = self.decoder(rep, edge_index, training=training)
632
+ x_init = ops.take(x, mask_nodes, axis=0)
633
+ x_rec = ops.take(recon, mask_nodes, axis=0)
634
+ loss_rec = criterion(x_rec, x_init, alpha=self.alpha_l) if self.loss_fn == "sce" else criterion(x_rec, x_init)
635
+
636
+ total_loss = loss_rec + self.lam * loss_latent
637
+
638
+ if epoch >= self.delayed_ema_epoch and training:
639
+ self.ema_update()
640
+
641
+ return total_loss
642
+
643
+ def call(self, x, edge_index, training: bool = False):
644
+ r"""Forward pass: returns node embeddings by default."""
645
+ return self.embed(x, edge_index, training=training)
646
+
647
+ def load_weights_from_checkpoint(
648
+ self,
649
+ checkpoint_path: Optional[str] = None,
650
+ dataset: Optional[str] = None,
651
+ folder: str = "checkpoints",
652
+ download: bool = True,
653
+ ):
654
+ r"""Loads weights from a PyTorch state dict checkpoint or Google Drive."""
655
+ return load_graphmae2_weights(
656
+ self,
657
+ checkpoint_path=checkpoint_path,
658
+ dataset=dataset,
659
+ folder=folder,
660
+ download=download,
661
+ )
662
+
663
+ @classmethod
664
+ def from_pretrained(
665
+ cls,
666
+ dataset: str = "ogbn-arxiv",
667
+ folder: str = "checkpoints",
668
+ download: bool = True,
669
+ **kwargs,
670
+ ) -> "GraphMAE2":
671
+ r"""Instantiates a GraphMAE2 model with pre-trained weights downloaded from Google Drive.
672
+
673
+ Args:
674
+ dataset (str): Dataset name (``"ogbn-arxiv"``, ``"ogbn-products"``,
675
+ ``"mag-scholar-f"``, or ``"ogbn-papers100M"``).
676
+ folder (str, optional): Directory to store/find checkpoints. (default: ``"checkpoints"``)
677
+ download (bool, optional): Whether to download checkpoint if missing locally. (default: ``True``)
678
+ **kwargs: Overrides for model hyperparameters.
679
+
680
+ Returns:
681
+ GraphMAE2: Model instance loaded with pre-trained weights.
682
+ """
683
+ key = _canonical_dataset_name(dataset)
684
+ if key not in GRAPHMAE2_PRETRAINED:
685
+ raise ValueError(
686
+ f"Unknown dataset '{dataset}'. Available pre-trained models: {list(GRAPHMAE2_PRETRAINED.keys())}"
687
+ )
688
+ cfg = dict(GRAPHMAE2_PRETRAINED[key])
689
+ cfg.pop("id")
690
+ cfg.pop("filename")
691
+ cfg.update(kwargs)
692
+
693
+ model = cls(**cfg)
694
+ load_graphmae2_weights(model, dataset=key, folder=folder, download=download)
695
+ return model
696
+
697
+ def __repr__(self) -> str:
698
+ return (
699
+ f"{self.__class__.__name__}(in_dim={self.in_dim}, "
700
+ f"num_hidden={self.num_hidden}, num_layers={self.num_layers}, "
701
+ f"num_dec_layers={self.num_dec_layers}, nhead={self.nhead})"
702
+ )
703
+
704
+
705
+ GRAPHMAE2_PRETRAINED = {
706
+ "ogbn-arxiv": {
707
+ "id": "1KdU5TbAg0lQwruO7SKenoFiC2MbaZQRr",
708
+ "filename": "gat_gat_1024_4_ogbn-arxiv_0.5_1024_checkpoint.pt",
709
+ "in_dim": 128,
710
+ "num_hidden": 1024,
711
+ "num_layers": 4,
712
+ "num_dec_layers": 1,
713
+ "nhead": 8,
714
+ "nhead_out": 1,
715
+ "activation": "prelu",
716
+ "norm": "layernorm",
717
+ "residual": True,
718
+ },
719
+ "ogbn-products": {
720
+ "id": "1Qk3bgK8H3bee3qmagH_hRJOfQmUD8PCW",
721
+ "filename": "gat_gat_1024_4_ogbn-products_0.5_1024_checkpoint.pt",
722
+ "in_dim": 100,
723
+ "num_hidden": 1024,
724
+ "num_layers": 4,
725
+ "num_dec_layers": 1,
726
+ "nhead": 4,
727
+ "nhead_out": 1,
728
+ "activation": "prelu",
729
+ "norm": "layernorm",
730
+ "residual": True,
731
+ },
732
+ "mag-scholar-f": {
733
+ "id": "1KpQk_OKbbo4qTLQYZ84pAJDy1sh4oZv2",
734
+ "filename": "gat_gat_1024_4_mag-scholar-f_0.5_1024_checkpoint.pt",
735
+ "in_dim": 128,
736
+ "num_hidden": 1024,
737
+ "num_layers": 4,
738
+ "num_dec_layers": 1,
739
+ "nhead": 8,
740
+ "nhead_out": 1,
741
+ "activation": "prelu",
742
+ "norm": "layernorm",
743
+ "residual": True,
744
+ },
745
+ "ogbn-papers100M": {
746
+ "id": "1zCD_vOckLfOXD1dWRY025A30QeuHsA_0",
747
+ "filename": "gat_gat_1024_4_ogbn-papers100M_0.5_1024_checkpoint.pt",
748
+ "in_dim": 128,
749
+ "num_hidden": 1024,
750
+ "num_layers": 4,
751
+ "num_dec_layers": 1,
752
+ "nhead": 8,
753
+ "nhead_out": 1,
754
+ "activation": "prelu",
755
+ "norm": "layernorm",
756
+ "residual": True,
757
+ },
758
+ }
759
+
760
+ DATASET_ALIASES = {
761
+ "arxiv": "ogbn-arxiv",
762
+ "products": "ogbn-products",
763
+ "mag": "mag-scholar-f",
764
+ "mag-scholar": "mag-scholar-f",
765
+ "papers100m": "ogbn-papers100M",
766
+ "ogbn-papers100m": "ogbn-papers100M",
767
+ "papers": "ogbn-papers100M",
768
+ }
769
+
770
+
771
+ def _canonical_dataset_name(name: Optional[str]) -> str:
772
+ if not name:
773
+ return ""
774
+ name_clean = name.strip()
775
+ if name_clean in GRAPHMAE2_PRETRAINED:
776
+ return name_clean
777
+ name_lower = name_clean.lower()
778
+ if name_lower in DATASET_ALIASES:
779
+ return DATASET_ALIASES[name_lower]
780
+ for k in GRAPHMAE2_PRETRAINED:
781
+ if k.lower() == name_lower:
782
+ return k
783
+ for k, v in GRAPHMAE2_PRETRAINED.items():
784
+ if v["filename"] == name_clean:
785
+ return k
786
+ return name_clean
787
+
788
+
789
+ def download_graphmae2_checkpoint(
790
+ dataset: str,
791
+ folder: str = "checkpoints",
792
+ log: bool = True,
793
+ ) -> str:
794
+ r"""Downloads a pre-trained GraphMAE2 checkpoint from Google Drive using download_google_url.
795
+
796
+ Google Drive folder: https://drive.google.com/drive/folders/1GiuP0PtIZaYlJWIrjvu73ZQCJGr6kGkh
797
+
798
+ Args:
799
+ dataset (str): Name of dataset (``"ogbn-arxiv"``, ``"ogbn-products"``,
800
+ ``"mag-scholar-f"``, or ``"ogbn-papers100M"``).
801
+ folder (str, optional): Target directory to save the checkpoint. (default: ``"checkpoints"``)
802
+ log (bool, optional): Whether to print download progress. (default: ``True``)
803
+
804
+ Returns:
805
+ str: Absolute path to the downloaded checkpoint file.
806
+ """
807
+ key = _canonical_dataset_name(dataset)
808
+ if key not in GRAPHMAE2_PRETRAINED:
809
+ raise ValueError(
810
+ f"Unknown dataset '{dataset}'. Available pre-trained checkpoints: {list(GRAPHMAE2_PRETRAINED.keys())}"
811
+ )
812
+ info = GRAPHMAE2_PRETRAINED[key]
813
+
814
+ target = os.path.join(folder, info["filename"])
815
+ if os.path.exists(target):
816
+ return target
817
+
818
+ # Check if local file exists in GraphMAE2-main/GraphMAE2_checkpoints when using default folder
819
+ if folder == "checkpoints":
820
+ local_alt = os.path.join("GraphMAE2-main", "GraphMAE2_checkpoints", info["filename"])
821
+ if os.path.exists(local_alt):
822
+ return local_alt
823
+
824
+ return download_google_url(
825
+ id=info["id"],
826
+ folder=folder,
827
+ filename=info["filename"],
828
+ log=log,
829
+ )
830
+
831
+
832
+ def load_graphmae2_weights(
833
+ model: GraphMAE2,
834
+ checkpoint_path: Optional[str] = None,
835
+ dataset: Optional[str] = None,
836
+ folder: str = "checkpoints",
837
+ download: bool = True,
838
+ ):
839
+ r"""Loads pre-trained weights from a GraphMAE2 PyTorch checkpoint (.pt).
840
+
841
+ If the checkpoint does not exist locally and download=True, it will be automatically
842
+ downloaded from the official Google Drive folder using download_google_url.
843
+
844
+ Args:
845
+ model (GraphMAE2): The target GraphMAE2 model instance.
846
+ checkpoint_path (str, optional): Local path to .pt file or dataset name.
847
+ dataset (str, optional): Dataset name if downloading from Google Drive.
848
+ folder (str, optional): Directory to store downloaded checkpoints. (default: ``"checkpoints"``)
849
+ download (bool, optional): Whether to download checkpoint if missing locally. (default: ``True``)
850
+ """
851
+ if checkpoint_path is None and dataset is None:
852
+ raise ValueError("Either checkpoint_path or dataset must be specified.")
853
+
854
+ path_to_load = checkpoint_path
855
+
856
+ candidate_dataset = dataset or (_canonical_dataset_name(checkpoint_path) if checkpoint_path else None)
857
+ if candidate_dataset in GRAPHMAE2_PRETRAINED:
858
+ if checkpoint_path and os.path.isfile(checkpoint_path):
859
+ path_to_load = checkpoint_path
860
+ else:
861
+ filename = GRAPHMAE2_PRETRAINED[candidate_dataset]["filename"]
862
+ alt_local = os.path.join("GraphMAE2-main", "GraphMAE2_checkpoints", filename)
863
+ default_local = os.path.join(folder, filename)
864
+ if os.path.isfile(alt_local):
865
+ path_to_load = alt_local
866
+ elif os.path.isfile(default_local):
867
+ path_to_load = default_local
868
+ elif download:
869
+ path_to_load = download_graphmae2_checkpoint(candidate_dataset, folder=folder)
870
+ else:
871
+ raise FileNotFoundError(f"Checkpoint for '{candidate_dataset}' not found at '{checkpoint_path}'.")
872
+ elif checkpoint_path and not os.path.isfile(checkpoint_path):
873
+ if download and dataset:
874
+ path_to_load = download_graphmae2_checkpoint(dataset, folder=folder)
875
+ else:
876
+ raise FileNotFoundError(f"Checkpoint file '{checkpoint_path}' not found.")
877
+
878
+ import torch
879
+
880
+ state_dict = torch.load(path_to_load, map_location="cpu")
881
+ if not isinstance(state_dict, dict):
882
+ raise ValueError(f"Expected a dict/state_dict in checkpoint, got {type(state_dict)}")
883
+
884
+ if not model.built:
885
+ model.build((None, model.in_dim))
886
+
887
+ def _to_tensor(t):
888
+ if hasattr(t, "detach"):
889
+ t = t.detach()
890
+ if hasattr(t, "numpy"):
891
+ t = t.numpy()
892
+ return ops.convert_to_tensor(np.array(t, dtype=np.float32), dtype="float32")
893
+
894
+ # 1. Mask tokens
895
+ if "enc_mask_token" in state_dict:
896
+ model.enc_mask_token.assign(_to_tensor(state_dict["enc_mask_token"]))
897
+ if "dec_mask_token" in state_dict:
898
+ model.dec_mask_token.assign(_to_tensor(state_dict["dec_mask_token"]))
899
+
900
+ # Helper for GAT module
901
+ def _load_gat(gat_module, prefix):
902
+ for i, layer in enumerate(gat_module.gat_layers):
903
+ p = f"{prefix}.gat_layers.{i}"
904
+ if f"{p}.fc.weight" in state_dict:
905
+ layer.fc.kernel.assign(_to_tensor(state_dict[f"{p}.fc.weight"].t()))
906
+ if f"{p}.attn_l" in state_dict:
907
+ layer.attn_l.assign(_to_tensor(state_dict[f"{p}.attn_l"]))
908
+ if f"{p}.attn_r" in state_dict:
909
+ layer.attn_r.assign(_to_tensor(state_dict[f"{p}.attn_r"]))
910
+ if f"{p}.bias" in state_dict and layer.bias is not None:
911
+ layer.bias.assign(_to_tensor(state_dict[f"{p}.bias"]))
912
+ if f"{p}.res_fc.weight" in state_dict and layer.res_fc is not None:
913
+ layer.res_fc.kernel.assign(_to_tensor(state_dict[f"{p}.res_fc.weight"].t()))
914
+ if f"{p}.activation.weight" in state_dict and hasattr(layer.activation, "alpha"):
915
+ layer.activation.alpha.assign(_to_tensor(state_dict[f"{p}.activation.weight"]))
916
+ if layer.norm is not None:
917
+ if f"{p}.norm.weight" in state_dict and hasattr(layer.norm, "gamma"):
918
+ layer.norm.gamma.assign(_to_tensor(state_dict[f"{p}.norm.weight"]))
919
+ if f"{p}.norm.bias" in state_dict and hasattr(layer.norm, "beta"):
920
+ layer.norm.beta.assign(_to_tensor(state_dict[f"{p}.norm.bias"]))
921
+
922
+ _load_gat(model.encoder, "encoder")
923
+ _load_gat(model.decoder, "decoder")
924
+ _load_gat(model.encoder_ema, "encoder_ema")
925
+
926
+ # 2. encoder_to_decoder
927
+ if "encoder_to_decoder.weight" in state_dict:
928
+ model.encoder_to_decoder.kernel.assign(_to_tensor(state_dict["encoder_to_decoder.weight"].t()))
929
+
930
+ # 3. Projectors
931
+ def _load_projector(proj_module, prefix):
932
+ if f"{prefix}.0.weight" in state_dict:
933
+ proj_module.layers[0].kernel.assign(_to_tensor(state_dict[f"{prefix}.0.weight"].t()))
934
+ if f"{prefix}.0.bias" in state_dict:
935
+ proj_module.layers[0].bias.assign(_to_tensor(state_dict[f"{prefix}.0.bias"]))
936
+ if f"{prefix}.1.weight" in state_dict and hasattr(proj_module.layers[1], "alpha"):
937
+ proj_module.layers[1].alpha.assign(_to_tensor(state_dict[f"{prefix}.1.weight"]))
938
+ if f"{prefix}.2.weight" in state_dict:
939
+ proj_module.layers[2].kernel.assign(_to_tensor(state_dict[f"{prefix}.2.weight"].t()))
940
+ if f"{prefix}.2.bias" in state_dict:
941
+ proj_module.layers[2].bias.assign(_to_tensor(state_dict[f"{prefix}.2.bias"]))
942
+
943
+ _load_projector(model.projector, "projector")
944
+ _load_projector(model.projector_ema, "projector_ema")
945
+
946
+ # 4. Predictor
947
+ if "predictor.0.weight" in state_dict and hasattr(model.predictor.layers[0], "alpha"):
948
+ model.predictor.layers[0].alpha.assign(_to_tensor(state_dict["predictor.0.weight"]))
949
+ if "predictor.1.weight" in state_dict:
950
+ model.predictor.layers[1].kernel.assign(_to_tensor(state_dict["predictor.1.weight"].t()))
951
+ if "predictor.1.bias" in state_dict:
952
+ model.predictor.layers[1].bias.assign(_to_tensor(state_dict["predictor.1.bias"]))
953
+
954
+ return model