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,101 @@
1
+ from keras import layers, ops
2
+
3
+ from k3_node.layers.conv.message_passing import MessagePassing
4
+ from k3_node.layers.conv.utils import gcn_norm, is_tracing
5
+
6
+
7
+ class SSGConv(MessagePassing):
8
+ r"""The simple spectral graph convolutional operator from the
9
+ `"Simple Spectral Graph Convolution" <https://arxiv.org/abs/2109.07191>`_ paper.
10
+
11
+ Args:
12
+ in_channels: Size of each input sample.
13
+ out_channels: Size of each output sample.
14
+ alpha: Teleport probability :math:`\alpha`.
15
+ K: Number of hops :math:`K`. (default: ``1``)
16
+ cached: If set to :obj:`True`, the layer will cache normalization coefficients.
17
+ (default: ``False``)
18
+ add_self_loops: If set to :obj:`False`, will not add self-loops.
19
+ (default: ``True``)
20
+ bias: If set to :obj:`False`, the layer will not learn an additive bias.
21
+ (default: ``True``)
22
+
23
+ Example:
24
+ ```python
25
+ import numpy as np
26
+ from k3_node.layers import SSGConv
27
+
28
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
29
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
30
+
31
+ layer = SSGConv(in_channels=8, out_channels=16, alpha=0.1, K=2)
32
+ out = layer(x, edge_index)
33
+ print(tuple(out.shape)) # (10, 16)
34
+ ```
35
+ """
36
+
37
+ weighted_sum_message = True
38
+
39
+ def __init__(
40
+ self,
41
+ in_channels: int,
42
+ out_channels: int,
43
+ alpha: float,
44
+ K: int = 1,
45
+ cached: bool = False,
46
+ add_self_loops: bool = True,
47
+ bias: bool = True,
48
+ **kwargs,
49
+ ):
50
+ super().__init__(aggr="add", **kwargs)
51
+ self.in_channels = in_channels
52
+ self.out_channels = out_channels
53
+ self.alpha = alpha
54
+ self.K = K
55
+ self.cached = cached
56
+ self.add_self_loops = add_self_loops
57
+ self.use_bias = bias
58
+
59
+ self.lin = layers.Dense(out_channels, use_bias=bias)
60
+ self._cached_edge_index = None
61
+ self._cached_norm = None
62
+
63
+ def build(self, input_shape):
64
+ feat_shape = input_shape[0] if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)) else input_shape
65
+ self.lin.build(feat_shape)
66
+ self.built = True
67
+
68
+ def call(self, x, edge_index=None, edge_weight=None, **kwargs):
69
+ if edge_index is None and isinstance(x, (tuple, list)):
70
+ x, edge_index = x[0], x[1]
71
+
72
+ if self.cached and self._cached_edge_index is not None:
73
+ edge_index = self._cached_edge_index
74
+ edge_weight = self._cached_norm
75
+ else:
76
+ num_nodes = x.shape[self.node_dim] if hasattr(x, "shape") and x.shape[self.node_dim] is not None else ops.shape(x)[self.node_dim]
77
+ edge_index, edge_weight = gcn_norm(
78
+ edge_index,
79
+ edge_weight,
80
+ num_nodes=num_nodes,
81
+ add_self_loops=self.add_self_loops,
82
+ flow=self.flow,
83
+ dtype=x.dtype,
84
+ )
85
+ if self.cached and not is_tracing(edge_index):
86
+ self._cached_edge_index = edge_index
87
+ self._cached_norm = edge_weight
88
+
89
+ out = self.alpha * x
90
+ h = x
91
+ for _ in range(self.K):
92
+ h = self.propagate(edge_index, x=h, edge_weight=edge_weight)
93
+ out = out + ((1.0 - self.alpha) / self.K) * h
94
+
95
+ return self.lin(out)
96
+
97
+ def message(self, x_j, edge_weight=None):
98
+ if edge_weight is None:
99
+ return x_j
100
+ return ops.expand_dims(edge_weight, -1) * x_j
101
+
@@ -0,0 +1,195 @@
1
+ import math
2
+ import keras
3
+ from keras import ops
4
+ from keras.layers import Dense, Dropout
5
+
6
+ from k3_node.layers.conv.message_passing import MessagePassing
7
+ from k3_node.layers.conv.utils import (
8
+ add_self_loops,
9
+ extend_mask_for_self_loops,
10
+ mask_edge_logits,
11
+ remove_self_loops_masked,
12
+ softmax,
13
+ )
14
+ from k3_node.ops.creation import full
15
+
16
+
17
+ class SuperGATConv(MessagePassing):
18
+ r"""The self-supervised graph attentional operator from the
19
+ `"How to Find Your Friendly Neighborhood: Graph Attention Design with Self-Supervision"
20
+ <https://openreview.net/forum?id=Wi5KUNlqWty>`_ paper.
21
+
22
+ Args:
23
+ attention_type (str): ``"MX"`` (mixed GO/DP) or ``"SD"`` (scaled dot-product).
24
+ neg_sample_ratio (float): Negative (random) pairs per positive edge in the attention loss.
25
+ edge_sample_ratio (float): Fraction of edges used as positives in the attention loss.
26
+ attention_loss_weight (float): If positive, the self-supervised attention loss is added to
27
+ the model loss (``model.losses``) with this weight while training, so ``fit`` optimizes
28
+ it automatically. PyG's example uses ``4.0``. (default: ``0.0``)
29
+
30
+ Example:
31
+ ```python
32
+ import numpy as np
33
+ from k3_node.layers import SuperGATConv
34
+
35
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
36
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
37
+
38
+ layer = SuperGATConv(in_channels=8, out_channels=16, heads=2)
39
+ out = layer(x, edge_index)
40
+ print(tuple(out.shape)) # (10, 32)
41
+ ```
42
+ """
43
+ def __init__(
44
+ self,
45
+ in_channels: int,
46
+ out_channels: int,
47
+ heads: int = 1,
48
+ concat: bool = True,
49
+ negative_slope: float = 0.2,
50
+ dropout: float = 0.0,
51
+ add_self_loops: bool = True,
52
+ bias: bool = True,
53
+ attention_type: str = "MX",
54
+ neg_sample_ratio: float = 0.5,
55
+ edge_sample_ratio: float = 1.0,
56
+ is_undirected: bool = False,
57
+ attention_loss_weight: float = 0.0,
58
+ **kwargs,
59
+ ):
60
+ kwargs.setdefault("aggr", "add")
61
+ super().__init__(node_dim=0, **kwargs)
62
+
63
+ assert attention_type in ["MX", "SD"]
64
+
65
+ self.in_channels = in_channels
66
+ self.out_channels = out_channels
67
+ self.heads = heads
68
+ self.concat = concat
69
+ self.negative_slope = negative_slope
70
+ self.dropout_rate = dropout
71
+ self.add_self_loops = add_self_loops
72
+ self.attention_type = attention_type
73
+ self.neg_sample_ratio = neg_sample_ratio
74
+ self.edge_sample_ratio = edge_sample_ratio
75
+ self.is_undirected = is_undirected
76
+ self.use_bias = bias
77
+
78
+ self.attention_loss_weight = attention_loss_weight
79
+ self.lin = Dense(heads * out_channels, use_bias=False)
80
+ self.dropout = Dropout(dropout)
81
+ self.seed_generator = keras.random.SeedGenerator()
82
+ self._last_attention_loss = None
83
+
84
+ if self.attention_type == "MX":
85
+ self.att_l = self.add_weight(
86
+ shape=(1, heads, out_channels),
87
+ initializer="glorot_uniform",
88
+ name="att_l",
89
+ )
90
+ self.att_r = self.add_weight(
91
+ shape=(1, heads, out_channels),
92
+ initializer="glorot_uniform",
93
+ name="att_r",
94
+ )
95
+ else:
96
+ self.att_l = None
97
+ self.att_r = None
98
+
99
+ if bias:
100
+ out_dim = heads * out_channels if concat else out_channels
101
+ self.bias = self.add_weight(
102
+ shape=(out_dim,),
103
+ initializer="zeros",
104
+ name="bias",
105
+ )
106
+ else:
107
+ self.bias = None
108
+
109
+ def build(self, input_shape=None):
110
+ self.lin.build((None, self.in_channels))
111
+ self.built = True
112
+
113
+ def call(self, inputs, edge_index=None, training=None, **kwargs):
114
+ if edge_index is None:
115
+ if isinstance(inputs, (list, tuple)) and len(inputs) == 2:
116
+ x, edge_index = inputs
117
+ else:
118
+ raise ValueError("Expected (x, edge_index) or x and edge_index")
119
+ else:
120
+ x = inputs
121
+
122
+ if not self.built:
123
+ self.build()
124
+
125
+ num_nodes = ops.shape(x)[0]
126
+ keep_mask = None
127
+ if self.add_self_loops:
128
+ edge_index, _, keep_mask = remove_self_loops_masked(edge_index)
129
+ edge_index, _ = add_self_loops(edge_index, num_nodes=num_nodes)
130
+ keep_mask = extend_mask_for_self_loops(keep_mask, num_nodes)
131
+
132
+ x = self.lin(x)
133
+ x = ops.reshape(x, (-1, self.heads, self.out_channels))
134
+
135
+ out = self.propagate(edge_index, x=x, keep_mask=keep_mask, training=training, size=(num_nodes, num_nodes))
136
+
137
+ if training:
138
+ loss = self._attention_loss(x, edge_index, num_nodes)
139
+ if self.attention_loss_weight:
140
+ self.add_loss(self.attention_loss_weight * loss)
141
+ from k3_node.layers.conv.utils import is_tracing
142
+ self._last_attention_loss = None if is_tracing(loss) else loss
143
+
144
+ if self.concat:
145
+ out = ops.reshape(out, (-1, self.heads * self.out_channels))
146
+ else:
147
+ out = ops.mean(out, axis=1)
148
+
149
+ if self.bias is not None:
150
+ out = out + self.bias
151
+
152
+ return out
153
+
154
+ def message(self, x_i, x_j, index=None, size_i=None, keep_mask=None, training=None):
155
+ if self.attention_type == "MX":
156
+ logits = ops.sum(x_i * x_j, axis=-1)
157
+ alpha = ops.sum(x_j * self.att_l, axis=-1) + ops.sum(x_i * self.att_r, axis=-1)
158
+ alpha = alpha * ops.sigmoid(logits)
159
+ else: # SD
160
+ alpha = ops.sum(x_i * x_j, axis=-1) / math.sqrt(self.out_channels)
161
+
162
+ alpha = ops.leaky_relu(alpha, negative_slope=self.negative_slope)
163
+ alpha = mask_edge_logits(alpha, keep_mask)
164
+ alpha = softmax(alpha, index, num_nodes=size_i, dim=0)
165
+ alpha = self.dropout(alpha, training=training)
166
+ return x_j * ops.expand_dims(alpha, -1)
167
+
168
+ def _attention_logits(self, x_i, x_j):
169
+ logits = ops.sum(x_i * x_j, axis=-1)
170
+ if self.attention_type == "SD":
171
+ logits = logits / math.sqrt(self.out_channels)
172
+ return logits
173
+
174
+ def _attention_loss(self, x, edge_index, num_nodes):
175
+ # Self-supervised attention loss: attention logits should separate real edges (label 1)
176
+ # from random node pairs (label 0). Edges are kept with probability `edge_sample_ratio` and
177
+ # random pairs with probability `neg_sample_ratio * edge_sample_ratio` (the expected counts
178
+ # used by PyG); weighting instead of slicing keeps all shapes static for XLA / jax.jit.
179
+ edge_index = ops.cast(edge_index, "int32")
180
+ num_edges = ops.shape(edge_index)[1]
181
+ neg = keras.random.randint(ops.shape(edge_index), 0, num_nodes, seed=self.seed_generator, dtype="int32")
182
+ pairs = ops.concatenate([edge_index, neg], axis=1)
183
+ logits = ops.mean(self._attention_logits(ops.take(x, pairs[1], axis=0), ops.take(x, pairs[0], axis=0)), axis=-1)
184
+ labels = ops.concatenate([ops.ones((num_edges,)), ops.zeros((num_edges,))])
185
+ keep = keras.random.uniform(ops.shape(labels), seed=self.seed_generator) < ops.concatenate([
186
+ full((num_edges,), self.edge_sample_ratio),
187
+ full((num_edges,), self.neg_sample_ratio * self.edge_sample_ratio),
188
+ ])
189
+ weights = ops.cast(keep, "float32")
190
+ losses = ops.binary_crossentropy(labels, logits, from_logits=True)
191
+ return ops.sum(losses * weights) / ops.maximum(ops.sum(weights), 1.0)
192
+
193
+ def get_attention_loss(self):
194
+ r"""The self-supervised attention loss of the last training call (as in PyG)."""
195
+ return self._last_attention_loss
@@ -0,0 +1,98 @@
1
+ from keras import layers, ops
2
+
3
+ from k3_node.layers.conv.message_passing import MessagePassing
4
+ from k3_node.layers.conv.utils import gcn_norm
5
+
6
+
7
+ class TAGConv(MessagePassing):
8
+ r"""The topology adaptive graph convolutional operator from the
9
+ `"Topology Adaptive Graph Convolutional Networks"
10
+ <https://arxiv.org/abs/1710.10370>`_ paper.
11
+
12
+ Args:
13
+ in_channels: Size of each input sample.
14
+ out_channels: Size of each output sample.
15
+ K: Number of hops :math:`K`. (default: ``3``)
16
+ bias: If set to :obj:`False`, the layer will not learn an additive bias.
17
+ (default: ``True``)
18
+ normalize: Whether to apply symmetric normalization. (default: ``True``)
19
+
20
+ Example:
21
+ ```python
22
+ import numpy as np
23
+ from k3_node.layers import TAGConv
24
+
25
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
26
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
27
+
28
+ layer = TAGConv(in_channels=8, out_channels=16, K=2)
29
+ out = layer(x, edge_index)
30
+ print(tuple(out.shape)) # (10, 16)
31
+ ```
32
+ """
33
+
34
+ weighted_sum_message = True
35
+
36
+ def __init__(
37
+ self,
38
+ in_channels: int,
39
+ out_channels: int,
40
+ K: int = 3,
41
+ bias: bool = True,
42
+ normalize: bool = True,
43
+ **kwargs,
44
+ ):
45
+ super().__init__(aggr="add", **kwargs)
46
+ self.in_channels = in_channels
47
+ self.out_channels = out_channels
48
+ self.K = K
49
+ self.normalize = normalize
50
+ self.use_bias = bias
51
+
52
+ self.lins = [layers.Dense(out_channels, use_bias=False) for _ in range(K + 1)]
53
+
54
+ def build(self, input_shape):
55
+ feat_shape = input_shape[0] if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)) else input_shape
56
+ for lin in self.lins:
57
+ lin.build(feat_shape)
58
+
59
+ if self.use_bias:
60
+ self.bias = self.add_weight(
61
+ shape=(self.out_channels,),
62
+ initializer="zeros",
63
+ name="bias",
64
+ )
65
+ else:
66
+ self.bias = None
67
+ self.built = True
68
+
69
+ def call(self, x, edge_index=None, edge_weight=None, **kwargs):
70
+ if edge_index is None and isinstance(x, (tuple, list)):
71
+ x, edge_index = x[0], x[1]
72
+
73
+ if self.normalize:
74
+ num_nodes = x.shape[self.node_dim] if hasattr(x, "shape") and x.shape[self.node_dim] is not None else ops.shape(x)[self.node_dim]
75
+ edge_index, edge_weight = gcn_norm(
76
+ edge_index,
77
+ edge_weight,
78
+ num_nodes=num_nodes,
79
+ add_self_loops=False,
80
+ flow=self.flow,
81
+ dtype=x.dtype,
82
+ )
83
+
84
+ out = self.lins[0](x)
85
+ h = x
86
+ for k in range(1, self.K + 1):
87
+ h = self.propagate(edge_index, x=h, edge_weight=edge_weight)
88
+ out = out + self.lins[k](h)
89
+
90
+ if self.bias is not None:
91
+ out = out + self.bias
92
+ return out
93
+
94
+ def message(self, x_j, edge_weight=None):
95
+ if edge_weight is None:
96
+ return x_j
97
+ return ops.expand_dims(edge_weight, -1) * x_j
98
+
@@ -0,0 +1,164 @@
1
+ """Cross-backend and compiled-mode consistency tests for convolution layers.
2
+
3
+ Each layer is run with identical seeded inputs and weights, and its eager output is
4
+ compared against golden outputs recorded with the torch backend (the backend that
5
+ ``tests_reference`` validates against PyG). On TensorFlow and JAX the output under
6
+ ``jit_compile=True`` must also match the eager output.
7
+
8
+ The input graph deliberately contains self-loops: layers that remove and re-add
9
+ self-loops must handle them without dynamic shapes when compiled.
10
+
11
+ To regenerate the golden file after an intentional numerical change::
12
+
13
+ KERAS_BACKEND=torch python k3_node/layers/conv/test_backend_consistency.py
14
+ """
15
+ import os
16
+ import os.path as osp
17
+
18
+ import numpy as np
19
+ import pytest
20
+ import keras
21
+ from keras import layers
22
+
23
+ import k3_node.layers as L
24
+
25
+ GOLDEN_PATH = osp.join(osp.dirname(__file__), "testdata", "backend_consistency_golden.npz")
26
+
27
+ N, E, C, O = 12, 30, 8, 5
28
+
29
+
30
+ def _graph():
31
+ rng = np.random.default_rng(0)
32
+ edge_index = rng.integers(0, N, size=(2, E))
33
+ loops = np.array([[0, 3, 7], [0, 3, 7]]) # guarantee pre-existing self-loops
34
+ return {
35
+ "x": rng.standard_normal((N, C)).astype("float32"),
36
+ "pos": rng.standard_normal((N, 3)).astype("float32"),
37
+ "normal": rng.standard_normal((N, 3)).astype("float32"),
38
+ "edge_index": np.concatenate([edge_index, loops], axis=1).astype("int32"),
39
+ }
40
+
41
+
42
+ def _mlp(units):
43
+ return keras.Sequential([layers.Dense(units, activation="relu"), layers.Dense(units)])
44
+
45
+
46
+ def _xe(layer, g):
47
+ return layer(g["x"], g["edge_index"])
48
+
49
+
50
+ # name -> (layer factory, call function)
51
+ CASES = {
52
+ "GCNConv": (lambda: L.GCNConv(C, O), _xe),
53
+ "GATConv": (lambda: L.GATConv(C, O, heads=2), _xe),
54
+ "GATv2Conv": (lambda: L.GATv2Conv(C, O, heads=2), _xe),
55
+ "SuperGATConv": (lambda: L.SuperGATConv(C, O, heads=2), _xe),
56
+ "ClusterGCNConv": (lambda: L.ClusterGCNConv(C, O), _xe),
57
+ "AGNNConv": (lambda: L.AGNNConv(), _xe),
58
+ "FeaStConv": (lambda: L.FeaStConv(C, O, heads=2), _xe),
59
+ "PointNetConv": (lambda: L.PointNetConv(local_nn=_mlp(O)), lambda l, g: l(g["x"], g["pos"], g["edge_index"])),
60
+ "PointTransformerConv": (
61
+ lambda: L.PointTransformerConv(C, O),
62
+ lambda l, g: l(g["x"], g["pos"], g["edge_index"]),
63
+ ),
64
+ "PPFConv": (
65
+ lambda: L.PPFConv(local_nn=_mlp(O)),
66
+ lambda l, g: l(g["x"], g["pos"], g["normal"], g["edge_index"]),
67
+ ),
68
+ "SAGEConv": (lambda: L.SAGEConv(C, O), _xe),
69
+ "GraphConv": (lambda: L.GraphConv(C, O), _xe),
70
+ "TransformerConv": (lambda: L.TransformerConv(C, O, heads=2), _xe),
71
+ "ChebConv": (lambda: L.ChebConv(C, O, K=3), _xe),
72
+ "TAGConv": (lambda: L.TAGConv(C, O, K=2), _xe),
73
+ "SGConv": (lambda: L.SGConv(C, O, K=2), _xe),
74
+ "LEConv": (lambda: L.LEConv(C, O), _xe),
75
+ "ResGatedGraphConv": (lambda: L.ResGatedGraphConv(C, O), _xe),
76
+ "GENConv": (lambda: L.GENConv(C, O), _xe),
77
+ "GINConv": (lambda: L.GINConv(_mlp(O)), _xe),
78
+ "EdgeConv": (lambda: L.EdgeConv(keras.Sequential([layers.Dense(O)])), _xe),
79
+ "MFConv": (lambda: L.MFConv(C, O), _xe),
80
+ "FiLMConv": (lambda: L.FiLMConv(C, O), _xe),
81
+ "GeneralConv": (lambda: L.GeneralConv(C, O), _xe),
82
+ "PNAConv": (
83
+ lambda: L.PNAConv(
84
+ C,
85
+ O,
86
+ aggregators=["mean", "max", "min", "std"],
87
+ scalers=["identity", "amplification"],
88
+ deg=np.array([0, 2, 4, 3, 2, 1], "int32"),
89
+ ),
90
+ _xe,
91
+ ),
92
+ }
93
+
94
+
95
+ class ConsistencyWrapper(keras.Model):
96
+ def __init__(self, layer, call_fn):
97
+ super().__init__()
98
+ self.layer = layer
99
+ self.call_fn = call_fn
100
+
101
+ def call(self, inputs):
102
+ out = self.call_fn(self.layer, inputs)
103
+ return out[0] if isinstance(out, (tuple, list)) else out
104
+
105
+
106
+ def _build(name):
107
+ factory, call_fn = CASES[name]
108
+ g = _graph()
109
+ model = ConsistencyWrapper(factory(), call_fn)
110
+ model(g)
111
+ rng = np.random.default_rng(123)
112
+ weights = []
113
+ for v in model.weights:
114
+ w = rng.standard_normal(v.shape) * 0.3
115
+ if "moving_variance" in v.path:
116
+ w = np.abs(w) + 0.5
117
+ weights.append(w.astype(v.dtype))
118
+ model.set_weights(weights)
119
+ shapes = np.array([str(tuple(v.shape)) for v in model.weights])
120
+ return model, g, shapes
121
+
122
+
123
+ def _eager(model, g):
124
+ return keras.ops.convert_to_numpy(model(g))
125
+
126
+
127
+ @pytest.fixture(scope="module")
128
+ def golden():
129
+ if not osp.exists(GOLDEN_PATH):
130
+ pytest.fail(f"Golden file missing; regenerate it (see module docstring): {GOLDEN_PATH}")
131
+ return np.load(GOLDEN_PATH)
132
+
133
+
134
+ @pytest.mark.parametrize("name", sorted(CASES))
135
+ def test_eager_matches_torch_golden(name, golden):
136
+ model, g, shapes = _build(name)
137
+ assert list(shapes) == list(golden[f"{name}/shapes"]), (
138
+ f"{name}: weight layout differs from the torch backend, so weights cannot be shared"
139
+ )
140
+ np.testing.assert_allclose(_eager(model, g), golden[f"{name}/out"], rtol=1e-4, atol=1e-5)
141
+
142
+
143
+ @pytest.mark.skipif(
144
+ keras.backend.backend() not in ("tensorflow", "jax"),
145
+ reason="jit_compile=True means XLA only on the TensorFlow and JAX backends",
146
+ )
147
+ @pytest.mark.parametrize("name", sorted(CASES))
148
+ def test_jit_matches_eager(name):
149
+ model, g, _ = _build(name)
150
+ expected = _eager(model, g)
151
+ model.compile(jit_compile=True)
152
+ np.testing.assert_allclose(np.asarray(model.predict_on_batch(g)), expected, rtol=1e-4, atol=1e-5)
153
+
154
+
155
+ if __name__ == "__main__":
156
+ assert keras.backend.backend() == "torch", "Golden outputs must be generated with KERAS_BACKEND=torch"
157
+ arrays = {}
158
+ for case in sorted(CASES):
159
+ m, graph, layer_shapes = _build(case)
160
+ arrays[f"{case}/out"] = _eager(m, graph)
161
+ arrays[f"{case}/shapes"] = layer_shapes
162
+ os.makedirs(osp.dirname(GOLDEN_PATH), exist_ok=True)
163
+ np.savez(GOLDEN_PATH, **arrays)
164
+ print(f"Wrote {len(CASES)} golden outputs to {GOLDEN_PATH}")