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,154 @@
1
+ from typing import Optional, Union, Tuple, List
2
+
3
+ from keras import layers, ops
4
+
5
+ from k3_node.layers.conv.message_passing import MessagePassing
6
+ from k3_node.layers.aggr.base import Aggregation
7
+
8
+
9
+ class SAGEConv(MessagePassing):
10
+ r"""The GraphSAGE operator from the `"Inductive Representation Learning on
11
+ Large Graphs" <https://arxiv.org/abs/1706.02216>`_ paper.
12
+
13
+ Args:
14
+ in_channels: Size of each input sample, or a tuple for bipartite graphs.
15
+ out_channels: Size of each output sample.
16
+ aggr: The aggregation scheme to use (``"mean"``, ``"max"``,
17
+ ``"lstm"``, etc.). (default: ``"mean"``)
18
+ normalize: If set to :obj:`True`, output features will be
19
+ :math:`\ell_2`-normalized. (default: :obj:`False`)
20
+ root_weight: If set to :obj:`False`, the layer will not add
21
+ the transformed root node features. (default: :obj:`True`)
22
+ project: If set to :obj:`True`, the layer will apply a linear
23
+ transformation followed by an activation to source node features.
24
+ (default: :obj:`False`)
25
+ bias: If set to :obj:`False`, the layer will not learn
26
+ an additive bias. (default: :obj:`True`)
27
+
28
+ Example:
29
+ ```python
30
+ import numpy as np
31
+ from k3_node.layers import SAGEConv
32
+
33
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
34
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
35
+
36
+ layer = SAGEConv(in_channels=8, out_channels=16)
37
+ out = layer(x, edge_index)
38
+ print(tuple(out.shape)) # (10, 16)
39
+ ```
40
+ """
41
+
42
+ def __init__(
43
+ self,
44
+ in_channels: Union[int, Tuple[int, int], None],
45
+ out_channels: Optional[int] = None,
46
+ aggr: Optional[Union[str, List[str], Aggregation]] = "mean",
47
+ normalize: bool = False,
48
+ root_weight: bool = True,
49
+ project: bool = False,
50
+ bias: bool = True,
51
+ **kwargs,
52
+ ):
53
+ if out_channels is None:
54
+ # Backward compatibility: SAGEConv(out_channels, ...)
55
+ out_channels = in_channels
56
+ in_channels = None
57
+
58
+ self.in_channels = in_channels
59
+ self.out_channels = out_channels
60
+ self.normalize = normalize
61
+ self.root_weight = root_weight
62
+ self.project = project
63
+ self.use_bias = bias
64
+
65
+ super().__init__(aggr=aggr, **kwargs)
66
+
67
+ if self.project:
68
+ self.lin_proj = layers.Dense(
69
+ in_channels[0] if isinstance(in_channels, (tuple, list)) else in_channels,
70
+ activation="relu",
71
+ use_bias=True,
72
+ )
73
+ else:
74
+ self.lin_proj = None
75
+
76
+ self.lin_l = layers.Dense(out_channels, use_bias=bias)
77
+ if self.root_weight:
78
+ self.lin_r = layers.Dense(out_channels, use_bias=False)
79
+ else:
80
+ self.lin_r = None
81
+
82
+ def build(self, input_shape):
83
+ if isinstance(input_shape, (tuple, list)) and len(input_shape) > 0 and isinstance(input_shape[0], (tuple, list)):
84
+ in_channels_l = input_shape[0][-1]
85
+ in_channels_r = input_shape[1][-1] if len(input_shape) > 1 and input_shape[1] is not None else in_channels_l
86
+ else:
87
+ in_channels_l = input_shape[-1]
88
+ in_channels_r = input_shape[-1]
89
+
90
+ self.lin_l.build((None, self.out_channels if self.project else in_channels_l))
91
+ if self.lin_r is not None:
92
+ self.lin_r.build((None, in_channels_r))
93
+ if self.lin_proj is not None:
94
+ self.lin_proj.build((None, in_channels_l))
95
+ self.built = True
96
+
97
+ def call(self, x, edge_index=None, size=None, **kwargs):
98
+ # Handle legacy calling: conv(x, adj) where adj is [N, N]
99
+ shape = getattr(edge_index, "shape", None)
100
+ if (
101
+ shape is not None
102
+ and len(shape) == 2
103
+ and shape[0] is not None
104
+ and shape[1] is not None
105
+ and shape[0] > 2
106
+ and shape[0] == shape[1]
107
+ ):
108
+ where_adj = ops.where(edge_index != 0)
109
+ where_adj = where_adj if not isinstance(where_adj, list) else where_adj
110
+ edge_index = ops.stack([where_adj[0], where_adj[1]], axis=0)
111
+
112
+ # Handle legacy calling: conv((x, adj))
113
+ if edge_index is None and isinstance(x, (tuple, list)) and len(x) == 2:
114
+ arg0, arg1 = x[0], x[1]
115
+ s1 = getattr(arg1, "shape", None)
116
+ if (
117
+ s1 is not None
118
+ and len(s1) == 2
119
+ and s1[0] is not None
120
+ and s1[1] is not None
121
+ and s1[0] > 2
122
+ and s1[0] == s1[1]
123
+ ):
124
+ where_adj = ops.where(arg1 != 0)
125
+ where_adj = where_adj if not isinstance(where_adj, list) else where_adj
126
+ edge_index = ops.stack([where_adj[0], where_adj[1]], axis=0)
127
+ x = arg0
128
+ elif s1 is not None and len(s1) >= 1 and s1[0] == 2:
129
+ edge_index = arg1
130
+ x = arg0
131
+
132
+ if not isinstance(x, (tuple, list)):
133
+ x_src = x
134
+ x_dst = x
135
+ else:
136
+ x_src, x_dst = x[0], x[1]
137
+
138
+ if self.project and self.lin_proj is not None:
139
+ x_src = self.lin_proj(x_src)
140
+
141
+ out = self.propagate(edge_index, x=(x_src, x_dst), size=size)
142
+ out = self.lin_l(out)
143
+
144
+ if self.root_weight and self.lin_r is not None and x_dst is not None:
145
+ out = out + self.lin_r(x_dst)
146
+
147
+ if self.normalize:
148
+ norm = ops.sqrt(ops.sum(ops.square(out), axis=-1, keepdims=True) + 1e-12)
149
+ out = out / norm
150
+
151
+ return out
152
+
153
+ def message(self, x_j):
154
+ return x_j
@@ -0,0 +1,96 @@
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 SGConv(MessagePassing):
8
+ r"""The simple graph convolutional operator from the `"Simplifying Graph
9
+ Convolutional Networks" <https://arxiv.org/abs/1902.07153>`_ paper.
10
+
11
+ Args:
12
+ in_channels: Size of each input sample.
13
+ out_channels: Size of each output sample.
14
+ K: Number of hops :math:`K`. (default: ``1``)
15
+ cached: If set to :obj:`True`, the layer will cache the computation of
16
+ :math:`\mathbf{\hat{D}}^{-1/2} \mathbf{\hat{A}} \mathbf{\hat{D}}^{-1/2}`.
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 SGConv
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 = SGConv(in_channels=8, out_channels=16, 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
+ K: int = 1,
44
+ cached: bool = False,
45
+ add_self_loops: bool = True,
46
+ bias: bool = True,
47
+ **kwargs,
48
+ ):
49
+ super().__init__(aggr="add", **kwargs)
50
+ self.in_channels = in_channels
51
+ self.out_channels = out_channels
52
+ self.K = K
53
+ self.cached = cached
54
+ self.add_self_loops = add_self_loops
55
+ self.use_bias = bias
56
+
57
+ self.lin = layers.Dense(out_channels, use_bias=bias)
58
+ self._cached_edge_index = None
59
+ self._cached_norm = None
60
+
61
+ def build(self, input_shape):
62
+ feat_shape = input_shape[0] if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)) else input_shape
63
+ self.lin.build(feat_shape)
64
+ self.built = True
65
+
66
+ def call(self, x, edge_index=None, edge_weight=None, **kwargs):
67
+ if edge_index is None and isinstance(x, (tuple, list)):
68
+ x, edge_index = x[0], x[1]
69
+
70
+ if self.cached and self._cached_edge_index is not None:
71
+ edge_index = self._cached_edge_index
72
+ edge_weight = self._cached_norm
73
+ else:
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=self.add_self_loops,
80
+ flow=self.flow,
81
+ dtype=x.dtype,
82
+ )
83
+ if self.cached and not is_tracing(edge_index):
84
+ self._cached_edge_index = edge_index
85
+ self._cached_norm = edge_weight
86
+
87
+ for _ in range(self.K):
88
+ x = self.propagate(edge_index, x=x, edge_weight=edge_weight)
89
+
90
+ return self.lin(x)
91
+
92
+ def message(self, x_j, edge_weight=None):
93
+ if edge_weight is None:
94
+ return x_j
95
+ return ops.expand_dims(edge_weight, -1) * x_j
96
+
@@ -0,0 +1,100 @@
1
+ from keras import ops
2
+ from keras.layers import Dense
3
+
4
+ from k3_node.layers.conv.message_passing import MessagePassing
5
+
6
+
7
+ class SignedConv(MessagePassing):
8
+ r"""The signed graph convolutional operator from the `"Signed Graph
9
+ Convolutional Network" <https://arxiv.org/abs/1808.06354>`_ paper.
10
+
11
+ Example:
12
+ ```python
13
+ import numpy as np
14
+ from k3_node.layers import SignedConv
15
+
16
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
17
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
18
+
19
+ pos_edge_index = edge_index[:, :15] # positive (e.g. "trust") edges
20
+ neg_edge_index = edge_index[:, 15:] # negative (e.g. "distrust") edges
21
+ layer = SignedConv(in_channels=8, out_channels=16, first_aggr=True)
22
+ out = layer(x, pos_edge_index, neg_edge_index)
23
+ print(tuple(out.shape)) # (10, 32): positive and negative embeddings, concatenated
24
+ ```
25
+ """
26
+ def __init__(
27
+ self,
28
+ in_channels: int,
29
+ out_channels: int,
30
+ first_aggr: bool,
31
+ bias: bool = True,
32
+ **kwargs,
33
+ ):
34
+ kwargs.setdefault("aggr", "mean")
35
+ super().__init__(**kwargs)
36
+
37
+ self.in_channels = in_channels
38
+ self.out_channels = out_channels
39
+ self.first_aggr = first_aggr
40
+ self.use_bias = bias
41
+
42
+ in_pos = in_channels if first_aggr else 2 * in_channels
43
+ self.lin_pos_l = Dense(out_channels, use_bias=False)
44
+ self.lin_pos_r = Dense(out_channels, use_bias=bias)
45
+ self.lin_neg_l = Dense(out_channels, use_bias=False)
46
+ self.lin_neg_r = Dense(out_channels, use_bias=bias)
47
+
48
+ def build(self, input_shape=None):
49
+ in_dim = self.in_channels if self.first_aggr else 2 * self.in_channels
50
+ self.lin_pos_l.build((None, in_dim))
51
+ self.lin_pos_r.build((None, self.in_channels))
52
+ self.lin_neg_l.build((None, in_dim))
53
+ self.lin_neg_r.build((None, self.in_channels))
54
+ self.built = True
55
+
56
+ def call(self, inputs, pos_edge_index=None, neg_edge_index=None, **kwargs):
57
+ if pos_edge_index is None:
58
+ if isinstance(inputs, (list, tuple)) and len(inputs) == 3:
59
+ x, pos_edge_index, neg_edge_index = inputs
60
+ else:
61
+ raise ValueError("Expected (x, pos_edge_index, neg_edge_index)")
62
+ else:
63
+ x = inputs
64
+
65
+ if not self.built:
66
+ self.build()
67
+
68
+ if isinstance(x, (list, tuple)):
69
+ x_src, x_dst = x
70
+ else:
71
+ x_src = x_dst = x
72
+
73
+ if self.first_aggr:
74
+ out_pos = self.propagate(pos_edge_index, x=(x_src, x_dst))
75
+ out_pos = self.lin_pos_l(out_pos) + self.lin_pos_r(x_dst)
76
+
77
+ out_neg = self.propagate(neg_edge_index, x=(x_src, x_dst))
78
+ out_neg = self.lin_neg_l(out_neg) + self.lin_neg_r(x_dst)
79
+
80
+ return ops.concatenate([out_pos, out_neg], axis=-1)
81
+ else:
82
+ F_in = self.in_channels
83
+ x_src_1, x_src_2 = x_src[..., :F_in], x_src[..., F_in:]
84
+ x_dst_1, x_dst_2 = x_dst[..., :F_in], x_dst[..., F_in:]
85
+
86
+ out_pos1 = self.propagate(pos_edge_index, x=(x_src_1, x_dst_1))
87
+ out_pos2 = self.propagate(neg_edge_index, x=(x_src_2, x_dst_2))
88
+ out_pos = ops.concatenate([out_pos1, out_pos2], axis=-1)
89
+ out_pos = self.lin_pos_l(out_pos) + self.lin_pos_r(x_dst_1)
90
+
91
+ out_neg1 = self.propagate(pos_edge_index, x=(x_src_2, x_dst_2))
92
+ out_neg2 = self.propagate(neg_edge_index, x=(x_src_1, x_dst_1))
93
+ out_neg = ops.concatenate([out_neg1, out_neg2], axis=-1)
94
+ out_neg = self.lin_neg_l(out_neg) + self.lin_neg_r(x_dst_2)
95
+
96
+ return ops.concatenate([out_pos, out_neg], axis=-1)
97
+
98
+ def message(self, x_j):
99
+ return x_j
100
+
@@ -0,0 +1,75 @@
1
+ from typing import Optional, Union, List
2
+ from keras import ops
3
+
4
+ from k3_node.layers.conv.message_passing import MessagePassing
5
+ from k3_node.layers.aggr.base import Aggregation
6
+
7
+
8
+ class SimpleConv(MessagePassing):
9
+ r"""A simple, parameter-free message passing operator.
10
+
11
+ Args:
12
+ aggr: The aggregation scheme to use (``"sum"``, ``"mean"``,
13
+ ``"min"``, ``"max"``, ``"mul"``). (default: ``"sum"``)
14
+ combine_root: The way to combine root node features with the
15
+ aggregated output (``"sum"``, ``"cat"``, ``"self_loop"``,
16
+ or :obj:`None`). (default: :obj:`None`)
17
+
18
+ Example:
19
+ ```python
20
+ import numpy as np
21
+ from k3_node.layers import SimpleConv
22
+
23
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
24
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
25
+
26
+ layer = SimpleConv(aggr="mean")
27
+ out = layer(x, edge_index)
28
+ print(tuple(out.shape)) # (10, 8)
29
+ ```
30
+ """
31
+
32
+ def __init__(
33
+ self,
34
+ aggr: Union[str, List[str], Aggregation, None] = "sum",
35
+ combine_root: Optional[str] = None,
36
+ **kwargs,
37
+ ):
38
+ super().__init__(aggr=aggr, **kwargs)
39
+ self.combine_root = combine_root
40
+ if combine_root is not None and combine_root not in ["sum", "cat", "self_loop"]:
41
+ raise ValueError(
42
+ f"combine_root must be 'sum', 'cat', 'self_loop', or None, got {combine_root}"
43
+ )
44
+
45
+ def build(self, input_shape):
46
+ self.built = True
47
+
48
+ def call(self, x, edge_index=None, edge_weight=None, size=None, **kwargs):
49
+ if edge_index is None and isinstance(x, (tuple, list)):
50
+ x, edge_index = x[0], x[1]
51
+
52
+ if not isinstance(x, (tuple, list)):
53
+ x_src, x_dst = x, x
54
+ else:
55
+ x_src, x_dst = x[0], x[1]
56
+
57
+ if self.combine_root == "self_loop":
58
+ from k3_node.layers.conv.utils import add_self_loops
59
+ num_nodes = ops.shape(x_src)[0]
60
+ edge_index, edge_weight = add_self_loops(edge_index, edge_weight, num_nodes=num_nodes)
61
+
62
+ out = self.propagate(edge_index, x=(x_src, x_dst), edge_weight=edge_weight, size=size)
63
+
64
+ if self.combine_root is not None and x_dst is not None and self.combine_root != "self_loop":
65
+ if self.combine_root == "sum":
66
+ out = out + x_dst
67
+ elif self.combine_root == "cat":
68
+ out = ops.concatenate([out, x_dst], axis=-1)
69
+
70
+ return out
71
+
72
+ def message(self, x_j, edge_weight=None):
73
+ if edge_weight is None:
74
+ return x_j
75
+ return ops.expand_dims(edge_weight, -1) * x_j
@@ -0,0 +1,182 @@
1
+ from typing import List, Tuple, Union
2
+ import keras
3
+ from keras import ops
4
+ from k3_node.layers.conv.message_passing import MessagePassing
5
+
6
+
7
+ class SplineConv(MessagePassing):
8
+ r"""The spline-based convolutional operator from the `"SplineCNN: Fast
9
+ Geometric Deep Learning with Continuous B-Spline Kernels"
10
+ <https://arxiv.org/abs/1711.08920>`_ paper.
11
+
12
+ Args:
13
+ in_channels (int or tuple): Size of each input sample.
14
+ out_channels (int): Size of each output sample.
15
+ dim (int): Pseudo-coordinate dimensionality.
16
+ kernel_size (int or List[int]): Size of the convolving kernel.
17
+ is_open_spline (bool or List[bool], optional): If set to :obj:`False`,
18
+ uses closed B-spline basis. (default: :obj:`True`)
19
+ degree (int, optional): B-spline basis degree. (default: :obj:`1`)
20
+ aggr (str, optional): The aggregation scheme to use (:obj:`"mean"`,
21
+ :obj:`"add"`, :obj:`"max"`). (default: :obj:`"mean"`)
22
+ root_weight (bool, optional): Whether to add transformed root node
23
+ features. (default: :obj:`True`)
24
+ bias (bool, optional): Whether to learn an additive bias. (default: :obj:`True`)
25
+
26
+ Example:
27
+ ```python
28
+ import numpy as np
29
+ from k3_node.layers import SplineConv
30
+
31
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
32
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
33
+ pseudo = np.random.rand(30, 2).astype("float32") # pseudo-coordinates in [0, 1]
34
+
35
+ layer = SplineConv(in_channels=8, out_channels=16, dim=2, kernel_size=3)
36
+ out = layer(x, edge_index, pseudo)
37
+ print(tuple(out.shape)) # (10, 16)
38
+ ```
39
+ """
40
+
41
+ def __init__(
42
+ self,
43
+ in_channels: Union[int, Tuple[int, int]],
44
+ out_channels: int,
45
+ dim: int,
46
+ kernel_size: Union[int, List[int]],
47
+ is_open_spline: Union[bool, List[bool]] = True,
48
+ degree: int = 1,
49
+ aggr: str = "mean",
50
+ root_weight: bool = True,
51
+ bias: bool = True,
52
+ **kwargs,
53
+ ):
54
+ super().__init__(aggr=aggr, **kwargs)
55
+
56
+ if isinstance(in_channels, int):
57
+ in_channels = (in_channels, in_channels)
58
+
59
+ if isinstance(kernel_size, int):
60
+ kernel_size = [kernel_size] * dim
61
+ assert len(kernel_size) == dim
62
+
63
+ self.in_channels = in_channels
64
+ self.out_channels = out_channels
65
+ self.dim = dim
66
+ self.kernel_size = kernel_size
67
+ self.degree = degree
68
+ self.root_weight = root_weight
69
+ self.use_bias = bias
70
+
71
+ # Total number of spline basis elements
72
+ k_prod = 1
73
+ for k in kernel_size:
74
+ k_prod *= k
75
+ self.K = k_prod
76
+
77
+ # Compute strides for multidimensional indexing
78
+ strides = []
79
+ stride = 1
80
+ for k in reversed(kernel_size):
81
+ strides.insert(0, stride)
82
+ stride *= k
83
+ self.strides = strides
84
+
85
+ def build(self, input_shape=None):
86
+ in_dim = self.in_channels[0]
87
+ if in_dim == -1 and input_shape is not None:
88
+ if isinstance(input_shape, (list, tuple)) and isinstance(input_shape[0], (list, tuple)):
89
+ in_dim = input_shape[0][-1]
90
+ elif isinstance(input_shape, (list, tuple)):
91
+ in_dim = input_shape[-1]
92
+ self.in_channels = (in_dim, self.in_channels[1] if self.in_channels[1] != -1 else in_dim)
93
+
94
+ self.weight = self.add_weight(
95
+ shape=(self.K, self.in_channels[0], self.out_channels),
96
+ initializer="glorot_uniform",
97
+ trainable=True,
98
+ name="weight",
99
+ )
100
+
101
+ if self.root_weight:
102
+ self.root_lin = keras.layers.Dense(self.out_channels, use_bias=False)
103
+ self.root_lin.build((None, self.in_channels[1]))
104
+
105
+ if self.use_bias:
106
+ self.bias = self.add_weight(
107
+ shape=(self.out_channels,),
108
+ initializer="zeros",
109
+ trainable=True,
110
+ name="bias",
111
+ )
112
+
113
+ super().build(input_shape)
114
+
115
+ def _spline_basis_1d(self, e_d, k_d):
116
+ # e_d: (E,) in [0, 1]
117
+ u = ops.clip(e_d * (k_d - 1), 0.0, float(k_d - 1))
118
+ i0 = ops.cast(ops.floor(u), "int32")
119
+ i0 = ops.clip(i0, 0, k_d - 1)
120
+ i1 = ops.clip(i0 + 1, 0, k_d - 1)
121
+ w1 = u - ops.cast(i0, u.dtype)
122
+ w0 = 1.0 - w1
123
+ return [(i0, w0), (i1, w1)]
124
+
125
+ def _compute_kernel(self, edge_attr):
126
+ # edge_attr: (E, D)
127
+ E = ops.shape(edge_attr)[0]
128
+ # Basis product over dimensions
129
+ dim_bases = [
130
+ self._spline_basis_1d(edge_attr[:, d], self.kernel_size[d])
131
+ for d in range(self.dim)
132
+ ]
133
+
134
+ # Cartesian product of basis across dimensions
135
+ basis_combinations = [([], 1.0)]
136
+ for d in range(self.dim):
137
+ new_combinations = []
138
+ stride = self.strides[d]
139
+ for curr_indices, curr_weight in basis_combinations:
140
+ for idx, w in dim_bases[d]:
141
+ new_idx = curr_indices + [idx * stride]
142
+ new_w = curr_weight * w
143
+ new_combinations.append((new_idx, new_w))
144
+ basis_combinations = new_combinations
145
+
146
+ # Construct edge-wise weight matrix W_eff: (E, C_in, C_out)
147
+ W_eff = ops.zeros((E, self.in_channels[0], self.out_channels), dtype=self.weight.dtype)
148
+ for idx_parts, basis_w in basis_combinations:
149
+ total_idx = idx_parts[0]
150
+ for part in idx_parts[1:]:
151
+ total_idx = total_idx + part
152
+ # total_idx: (E,)
153
+ basis_weights = ops.take(self.weight, total_idx, axis=0) # (E, C_in, C_out)
154
+ basis_w_expanded = ops.expand_dims(ops.expand_dims(basis_w, axis=-1), axis=-1)
155
+ W_eff = W_eff + basis_w_expanded * basis_weights
156
+
157
+ return W_eff
158
+
159
+ def call(self, x, edge_index, edge_attr=None, **kwargs):
160
+ if isinstance(x, (list, tuple)):
161
+ x_src, x_dst = x[0], x[1]
162
+ else:
163
+ x_src = x_dst = x
164
+
165
+ out = self.propagate(edge_index, x=(x_src, x_dst), edge_attr=edge_attr)
166
+
167
+ if self.root_weight and x_dst is not None:
168
+ out = out + self.root_lin(x_dst)
169
+
170
+ if self.use_bias:
171
+ out = out + self.bias
172
+
173
+ return out
174
+
175
+ def message(self, x_j, edge_attr):
176
+ if edge_attr is None:
177
+ return x_j
178
+ W_eff = self._compute_kernel(edge_attr) # (E, C_in, C_out)
179
+ # x_j: (E, C_in)
180
+ # m = sum_c x_j * W_eff
181
+ return ops.einsum("ec,ecd->ed", x_j, W_eff)
182
+