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,1156 @@
1
+ import math
2
+ import os
3
+ from typing import Optional, Union, Tuple, List, Dict, Any, Callable
4
+
5
+ import numpy as np
6
+ import keras
7
+ from keras import layers, ops
8
+
9
+ from k3_node.layers.attention.pair_attention import (
10
+ SelfMultiheadAttentionWithPair,
11
+ TransformerEncoderLayerWithPair,
12
+ _get_activation,
13
+ )
14
+ from k3_node.data.download import download_url
15
+
16
+
17
+ # ==============================================================================
18
+ # Pretrained Weights Registry
19
+ # ==============================================================================
20
+
21
+ UNIMOL_PRETRAINED_URLS: Dict[str, Dict[str, str]] = {
22
+ "mol_pre_no_h": {
23
+ "url": "https://github.com/deepmodeling/Uni-Mol/releases/download/v0.1/mol_pre_no_h_220816.pt",
24
+ "filename": "mol_pre_no_h_220816.pt",
25
+ "dict": "mol.dict.txt",
26
+ "dict_url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/mol.dict.txt",
27
+ "description": "Uni-Mol molecular pretraining model (no hydrogen)",
28
+ },
29
+ "mol_pre_all_h": {
30
+ "url": "https://github.com/deepmodeling/Uni-Mol/releases/download/v0.1/mol_pre_all_h_220816.pt",
31
+ "filename": "mol_pre_all_h_220816.pt",
32
+ "dict": "mol.dict.txt",
33
+ "dict_url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/mol.dict.txt",
34
+ "description": "Uni-Mol molecular pretraining model (all hydrogen)",
35
+ },
36
+ "pocket_pre": {
37
+ "url": "https://github.com/deepmodeling/Uni-Mol/releases/download/v0.1/pocket_pre_220816.pt",
38
+ "filename": "pocket_pre_220816.pt",
39
+ "dict": "poc.dict.txt",
40
+ "dict_url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/poc.dict.txt",
41
+ "description": "Uni-Mol candidate protein pocket pretraining model",
42
+ },
43
+ "mp_all_h": {
44
+ "url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/mp_all_h_230313.pt",
45
+ "filename": "mp_all_h_230313.pt",
46
+ "dict": "mp.dict.txt",
47
+ "dict_url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/mp.dict.txt",
48
+ "description": "Uni-Mol crystal Materials Project pretraining model",
49
+ },
50
+ "oled_pre_no_h": {
51
+ "url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/oled_pre_no_h_230101.pt",
52
+ "filename": "oled_pre_no_h_230101.pt",
53
+ "dict": "oled.dict.txt",
54
+ "dict_url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/oled.dict.txt",
55
+ "description": "Uni-Mol OLED molecule pretraining model",
56
+ },
57
+ "qm9": {
58
+ "url": "https://github.com/deepmodeling/Uni-Mol/releases/download/v0.1/qm9_220908.pt",
59
+ "filename": "qm9_220908.pt",
60
+ "dict": "mol.dict.txt",
61
+ "dict_url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/mol.dict.txt",
62
+ "description": "Uni-Mol conformation generation fine-tuned on QM9",
63
+ },
64
+ "drugs": {
65
+ "url": "https://github.com/deepmodeling/Uni-Mol/releases/download/v0.1/drugs_220908.pt",
66
+ "filename": "drugs_220908.pt",
67
+ "dict": "mol.dict.txt",
68
+ "dict_url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/mol.dict.txt",
69
+ "description": "Uni-Mol conformation generation fine-tuned on GEOM-Drugs",
70
+ },
71
+ "binding_pose": {
72
+ "url": "https://github.com/deepmodeling/Uni-Mol/releases/download/v0.1/binding_pose_220908.pt",
73
+ "filename": "binding_pose_220908.pt",
74
+ "dict": "mol.dict.txt",
75
+ "dict_url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/mol.dict.txt",
76
+ "description": "Uni-Mol protein-ligand binding pose prediction",
77
+ },
78
+ }
79
+
80
+ UNIMOL_ALIASES = {
81
+ "molecule": "mol_pre_no_h",
82
+ "molecule_no_h": "mol_pre_no_h",
83
+ "molecule_all_h": "mol_pre_all_h",
84
+ "protein": "pocket_pre",
85
+ "pocket": "pocket_pre",
86
+ "poc_pre": "pocket_pre",
87
+ "crystal": "mp_all_h",
88
+ "mp": "mp_all_h",
89
+ "oled": "oled_pre_no_h",
90
+ }
91
+
92
+
93
+ # ==============================================================================
94
+ # Layers
95
+ # ==============================================================================
96
+
97
+ class GaussianLayer(layers.Layer):
98
+ r"""Gaussian basis function (GBF) expansion over pairwise distances modulated by edge types.
99
+
100
+ Args:
101
+ num_kernel (int, optional): Number of Gaussian kernels. (default: ``128``)
102
+ edge_types (int, optional): Number of distinct pairwise edge types. (default: ``1024``)
103
+ **kwargs: Additional layer arguments.
104
+
105
+ Example:
106
+ ```python
107
+ import numpy as np
108
+ from k3_node.models import UniMolGaussianLayer
109
+
110
+ dist = np.random.rand(2, 5, 5).astype("float32") * 5.0 # pairwise atom distances
111
+ edge_type = np.random.randint(0, 128, size=(2, 5, 5)) # atom-pair type ids
112
+
113
+ layer = UniMolGaussianLayer(num_kernel=32, edge_types=128)
114
+ print(tuple(layer(dist, edge_type).shape)) # (2, 5, 5, 32): Gaussian distance features per atom pair
115
+ ```
116
+ """
117
+
118
+ def __init__(self, num_kernel: int = 128, edge_types: int = 1024, **kwargs):
119
+ super().__init__(**kwargs)
120
+ self.num_kernel = num_kernel
121
+ self.edge_types = edge_types
122
+
123
+ self.means = layers.Embedding(1, num_kernel, embeddings_initializer="uniform", name="means")
124
+ self.stds = layers.Embedding(1, num_kernel, embeddings_initializer="uniform", name="stds")
125
+ self.mul = layers.Embedding(edge_types, 1, embeddings_initializer="ones", name="mul")
126
+ self.bias = layers.Embedding(edge_types, 1, embeddings_initializer="zeros", name="bias")
127
+
128
+ def build(self, input_shape=None):
129
+ if not self.built:
130
+ self.means.build(None)
131
+ self.stds.build(None)
132
+ self.mul.build(None)
133
+ self.bias.build(None)
134
+ super().build(input_shape)
135
+
136
+ def call(self, dist, edge_type):
137
+ r"""
138
+ Args:
139
+ dist (Tensor): Pairwise distance matrix of shape ``[batch_size, seq_len, seq_len]``.
140
+ edge_type (Tensor): Pairwise edge type indices of shape ``[batch_size, seq_len, seq_len]``.
141
+
142
+ Returns:
143
+ Tensor: Gaussian basis expansion of shape ``[batch_size, seq_len, seq_len, num_kernel]``.
144
+ """
145
+ mul = ops.cast(self.mul(edge_type), dist.dtype)
146
+ bias = ops.cast(self.bias(edge_type), dist.dtype)
147
+
148
+ # x: [B, N, N, 1]
149
+ x = mul * ops.expand_dims(dist, axis=-1) + bias
150
+
151
+ # Means and stds: [num_kernel]
152
+ zero_idx = ops.zeros((1,), dtype="int32")
153
+ mean = ops.reshape(self.means(zero_idx), (-1,))
154
+ std = ops.abs(ops.reshape(self.stds(zero_idx), (-1,))) + 1e-5
155
+
156
+ a = math.sqrt(2.0 * math.pi)
157
+ diff = (x - mean) / std
158
+ return ops.exp(-0.5 * ops.power(diff, 2)) / (a * std)
159
+
160
+
161
+ class NumericalEmbed(layers.Layer):
162
+ r"""Numerical embedding layer for continuous edge features.
163
+
164
+ Example:
165
+ ```python
166
+ import numpy as np
167
+ from k3_node.models import UniMolNumericalEmbed
168
+
169
+ dist = np.random.rand(2, 5, 5).astype("float32") * 5.0 # pairwise atom distances
170
+ edge_type = np.random.randint(0, 128, size=(2, 5, 5)) # atom-pair type ids
171
+
172
+ layer = UniMolNumericalEmbed(num_kernel=32, edge_types=128)
173
+ print(tuple(layer(dist, edge_type).shape)) # (2, 5, 5, 32)
174
+ ```
175
+ """
176
+
177
+ def __init__(self, num_kernel: int = 128, edge_types: int = 1024, activation_fn: str = "gelu", **kwargs):
178
+ super().__init__(**kwargs)
179
+ self.num_kernel = num_kernel
180
+ self.edge_types = edge_types
181
+ self.mul = layers.Embedding(edge_types, 1, name="mul")
182
+ self.bias = layers.Embedding(edge_types, 1, name="bias")
183
+ self.w_edge = layers.Embedding(edge_types, num_kernel, name="w_edge")
184
+ self.proj = NonLinearHead(1, num_kernel, activation_fn=activation_fn, hidden=2 * num_kernel, name="proj")
185
+ self.ln = layers.LayerNormalization(axis=-1, epsilon=1e-5, name="ln")
186
+
187
+ def build(self, input_shape=None):
188
+ if not self.built:
189
+ self.mul.build(None)
190
+ self.bias.build(None)
191
+ self.w_edge.build(None)
192
+ self.proj.build((None, None, None, 1))
193
+ self.ln.build((None, None, None, self.num_kernel))
194
+ super().build(input_shape)
195
+
196
+ def call(self, dist, edge_type):
197
+ mul = ops.cast(self.mul(edge_type), dist.dtype)
198
+ bias = ops.cast(self.bias(edge_type), dist.dtype)
199
+ w_edge = ops.cast(self.w_edge(edge_type), dist.dtype)
200
+
201
+ edge_feat = mul * ops.expand_dims(dist, axis=-1) + bias
202
+ edge_feat = self.proj(edge_feat)
203
+ edge_feat = edge_feat + w_edge
204
+ return self.ln(edge_feat)
205
+
206
+
207
+ class NonLinearHead(layers.Layer):
208
+ r"""Two-layer feed-forward network with activation for feature projection.
209
+
210
+ Example:
211
+ ```python
212
+ import numpy as np
213
+ from k3_node.models import UniMolNonLinearHead
214
+
215
+ x = np.random.rand(2, 5, 32).astype("float32") # [batch, atoms, embed_dim]
216
+
217
+ head = UniMolNonLinearHead(input_dim=32, out_dim=16, activation_fn="gelu")
218
+ print(tuple(head(x).shape)) # (2, 5, 16)
219
+ ```
220
+ """
221
+
222
+ def __init__(
223
+ self,
224
+ input_dim: int,
225
+ out_dim: int,
226
+ activation_fn: Union[str, Callable] = "gelu",
227
+ hidden: Optional[int] = None,
228
+ **kwargs,
229
+ ):
230
+ super().__init__(**kwargs)
231
+ self.input_dim = input_dim
232
+ self.out_dim = out_dim
233
+ self.hidden = hidden or input_dim
234
+ self.activation_fn_name = activation_fn
235
+
236
+ self.linear1 = layers.Dense(self.hidden, name="linear1")
237
+ self.act = _get_activation(activation_fn)
238
+ self.linear2 = layers.Dense(out_dim, name="linear2")
239
+
240
+ def build(self, input_shape=None):
241
+ if not self.built:
242
+ self.linear1.build((None, None, None, self.input_dim) if len(input_shape or ()) == 4 else (None, None, self.input_dim))
243
+ self.linear2.build((None, None, None, self.hidden) if len(input_shape or ()) == 4 else (None, None, self.hidden))
244
+ super().build(input_shape)
245
+
246
+ def call(self, x):
247
+ x = self.linear1(x)
248
+ if self.act is not None:
249
+ x = self.act(x)
250
+ x = self.linear2(x)
251
+ return x
252
+
253
+
254
+ class DistanceHead(layers.Layer):
255
+ r"""Symmetrized distance prediction head from pair representations.
256
+
257
+ Example:
258
+ ```python
259
+ import numpy as np
260
+ from k3_node.models import UniMolDistanceHead
261
+
262
+ pair = np.random.rand(2, 6, 6, 8).astype("float32") # pair representation with 8 heads
263
+ print(tuple(UniMolDistanceHead(heads=8)(pair).shape)) # (2, 6, 6): predicted distance matrix
264
+ ```
265
+ """
266
+
267
+ def __init__(self, heads: int, activation_fn: Union[str, Callable] = "gelu", **kwargs):
268
+ super().__init__(**kwargs)
269
+ self.heads = heads
270
+ self.dense = layers.Dense(heads, name="dense")
271
+ self.act = _get_activation(activation_fn)
272
+ self.layer_norm = layers.LayerNormalization(axis=-1, epsilon=1e-5, name="layer_norm")
273
+ self.out_proj = layers.Dense(1, name="out_proj")
274
+
275
+ def build(self, input_shape=None):
276
+ if not self.built:
277
+ self.dense.build((None, None, None, self.heads))
278
+ self.layer_norm.build((None, None, None, self.heads))
279
+ self.out_proj.build((None, None, None, self.heads))
280
+ super().build(input_shape)
281
+
282
+ def call(self, x):
283
+ r"""
284
+ Args:
285
+ x (Tensor): Pair tensor of shape ``[batch_size, seq_len, seq_len, heads]``.
286
+
287
+ Returns:
288
+ Tensor: Symmetrized predicted distances ``[batch_size, seq_len, seq_len]``.
289
+ """
290
+ x = self.dense(x)
291
+ if self.act is not None:
292
+ x = self.act(x)
293
+ x = self.layer_norm(x)
294
+ x = ops.squeeze(self.out_proj(x), axis=-1) # [B, N, N]
295
+ return 0.5 * (x + ops.transpose(x, (0, 2, 1)))
296
+
297
+
298
+ class LinearHead(layers.Layer):
299
+ r"""Linear classification/regression head.
300
+
301
+ Example:
302
+ ```python
303
+ import numpy as np
304
+ from k3_node.models import UniMolLinearHead
305
+
306
+ x = np.random.rand(2, 5, 32).astype("float32") # [batch, atoms, embed_dim]
307
+
308
+ head = UniMolLinearHead(input_dim=32, num_classes=3)
309
+ print(tuple(head(x).shape)) # (2, 5, 3)
310
+ ```
311
+ """
312
+
313
+ def __init__(self, input_dim: int, num_classes: int, pooler_dropout: float = 0.0, **kwargs):
314
+ super().__init__(**kwargs)
315
+ self.input_dim = input_dim
316
+ self.num_classes = num_classes
317
+ self.pooler_dropout = pooler_dropout
318
+
319
+ self.dropout = layers.Dropout(pooler_dropout) if pooler_dropout > 0.0 else None
320
+ self.out_proj = layers.Dense(num_classes, name="out_proj")
321
+
322
+ def build(self, input_shape=None):
323
+ if not self.built:
324
+ self.out_proj.build((None, self.input_dim))
325
+ super().build(input_shape)
326
+
327
+ def call(self, features, training: bool = False):
328
+ x = features
329
+ if self.dropout is not None:
330
+ x = self.dropout(x, training=training)
331
+ return self.out_proj(x)
332
+
333
+
334
+ class ClassificationHead(layers.Layer):
335
+ r"""Two-layer sentence/graph-level classification head.
336
+
337
+ Example:
338
+ ```python
339
+ import numpy as np
340
+ from k3_node.models import UniMolClassificationHead
341
+
342
+ x = np.random.rand(2, 5, 32).astype("float32") # [batch, atoms, embed_dim]
343
+
344
+ head = UniMolClassificationHead(input_dim=32, inner_dim=32, num_classes=3)
345
+ print(tuple(head(x).shape)) # (2, 3): classifies from the first ([CLS]) token
346
+ ```
347
+ """
348
+
349
+ def __init__(
350
+ self,
351
+ input_dim: int,
352
+ inner_dim: int,
353
+ num_classes: int,
354
+ activation_fn: Union[str, Callable] = "gelu",
355
+ pooler_dropout: float = 0.0,
356
+ **kwargs,
357
+ ):
358
+ super().__init__(**kwargs)
359
+ self.input_dim = input_dim
360
+ self.inner_dim = inner_dim
361
+ self.num_classes = num_classes
362
+ self.pooler_dropout = pooler_dropout
363
+
364
+ self.dense = layers.Dense(inner_dim, name="dense")
365
+ self.act = _get_activation(activation_fn)
366
+ self.dropout = layers.Dropout(pooler_dropout) if pooler_dropout > 0.0 else None
367
+ self.out_proj = layers.Dense(num_classes, name="out_proj")
368
+
369
+ def build(self, input_shape=None):
370
+ if not self.built:
371
+ self.dense.build((None, self.input_dim))
372
+ self.out_proj.build((None, self.inner_dim))
373
+ super().build(input_shape)
374
+
375
+ def call(self, features, training: bool = False):
376
+ x = features
377
+ if len(ops.shape(features)) == 3:
378
+ x = features[:, 0, :] # CLS token
379
+ if self.dropout is not None:
380
+ x = self.dropout(x, training=training)
381
+ x = self.dense(x)
382
+ if self.act is not None:
383
+ x = self.act(x)
384
+ if self.dropout is not None:
385
+ x = self.dropout(x, training=training)
386
+ return self.out_proj(x)
387
+
388
+
389
+ class MaskLMHead(layers.Layer):
390
+ r"""Masked language modeling head for predicting masked atom tokens.
391
+
392
+ Example:
393
+ ```python
394
+ import numpy as np
395
+ from k3_node.models import UniMolMaskLMHead
396
+
397
+ x = np.random.rand(2, 5, 32).astype("float32") # [batch, atoms, embed_dim]
398
+
399
+ head = UniMolMaskLMHead(embed_dim=32, output_dim=64) # logits over a 64-token vocabulary
400
+ print(tuple(head(x).shape)) # (2, 5, 64)
401
+ ```
402
+ """
403
+
404
+ def __init__(
405
+ self,
406
+ embed_dim: int,
407
+ output_dim: int,
408
+ activation_fn: Union[str, Callable] = "gelu",
409
+ **kwargs,
410
+ ):
411
+ super().__init__(**kwargs)
412
+ self.embed_dim = embed_dim
413
+ self.output_dim = output_dim
414
+ self.dense = layers.Dense(embed_dim, name="dense")
415
+ self.act = _get_activation(activation_fn)
416
+ self.layer_norm = layers.LayerNormalization(axis=-1, epsilon=1e-5, name="layer_norm")
417
+ self.out_proj = layers.Dense(output_dim, name="out_proj")
418
+
419
+ def build(self, input_shape=None):
420
+ if not self.built:
421
+ self.dense.build((None, None, self.embed_dim))
422
+ self.layer_norm.build((None, None, self.embed_dim))
423
+ self.out_proj.build((None, None, self.embed_dim))
424
+ super().build(input_shape)
425
+
426
+ def call(self, features, masked_tokens=None):
427
+ x = self.dense(features)
428
+ if self.act is not None:
429
+ x = self.act(x)
430
+ x = self.layer_norm(x)
431
+ x = self.out_proj(x)
432
+ return x
433
+
434
+
435
+ # ==============================================================================
436
+ # Backbone Transformer Encoder
437
+ # ==============================================================================
438
+
439
+ class UniMolTransformerEncoder(layers.Layer):
440
+ r"""Transformer Encoder backbone for Uni-Mol with pair attention bias propagation."""
441
+
442
+ def __init__(
443
+ self,
444
+ encoder_layers: int = 15,
445
+ embed_dim: int = 512,
446
+ ffn_embed_dim: int = 2048,
447
+ attention_heads: int = 64,
448
+ emb_dropout: float = 0.1,
449
+ dropout: float = 0.1,
450
+ attention_dropout: float = 0.1,
451
+ activation_dropout: float = 0.0,
452
+ max_seq_len: int = 512,
453
+ activation_fn: Union[str, Callable] = "gelu",
454
+ post_ln: bool = False,
455
+ no_final_head_layer_norm: bool = False,
456
+ **kwargs,
457
+ ):
458
+ super().__init__(**kwargs)
459
+ self.num_layers = encoder_layers
460
+ self.embed_dim = embed_dim
461
+ self.ffn_embed_dim = ffn_embed_dim
462
+ self.attention_heads = attention_heads
463
+ self.emb_dropout_rate = emb_dropout
464
+ self.post_ln = post_ln
465
+
466
+ self.emb_layer_norm = layers.LayerNormalization(axis=-1, epsilon=1e-5, name="emb_layer_norm")
467
+ self.emb_dropout = layers.Dropout(emb_dropout) if emb_dropout > 0.0 else None
468
+
469
+ self.final_layer_norm = None if post_ln else layers.LayerNormalization(axis=-1, epsilon=1e-5, name="final_layer_norm")
470
+ self.final_head_layer_norm = None if no_final_head_layer_norm else layers.LayerNormalization(axis=-1, epsilon=1e-5, name="final_head_layer_norm")
471
+
472
+ self.layers_list = [
473
+ TransformerEncoderLayerWithPair(
474
+ embed_dim=embed_dim,
475
+ ffn_embed_dim=ffn_embed_dim,
476
+ attention_heads=attention_heads,
477
+ dropout=dropout,
478
+ attention_dropout=attention_dropout,
479
+ activation_dropout=activation_dropout,
480
+ activation_fn=activation_fn,
481
+ post_ln=post_ln,
482
+ name=f"layer_{i}",
483
+ )
484
+ for i in range(encoder_layers)
485
+ ]
486
+
487
+ def build(self, input_shape=None):
488
+ if not self.built:
489
+ self.emb_layer_norm.build((None, None, self.embed_dim))
490
+ if self.final_layer_norm is not None:
491
+ self.final_layer_norm.build((None, None, self.embed_dim))
492
+ if self.final_head_layer_norm is not None:
493
+ self.final_head_layer_norm.build((None, None, None, self.attention_heads))
494
+ for layer in self.layers_list:
495
+ layer.build((None, None, self.embed_dim))
496
+ super().build(input_shape)
497
+
498
+ def call(self, emb, attn_mask=None, padding_mask=None, training: bool = False):
499
+ shape = ops.shape(emb)
500
+ bsz = shape[0]
501
+ seq_len = shape[1]
502
+
503
+ x = self.emb_layer_norm(emb)
504
+ if self.emb_dropout is not None:
505
+ x = self.emb_dropout(x, training=training)
506
+
507
+ if padding_mask is not None:
508
+ mask_expanded = ops.expand_dims(ops.cast(padding_mask, x.dtype), axis=-1)
509
+ x = x * (1.0 - mask_expanded)
510
+
511
+ if attn_mask is not None:
512
+ mask_shape = ops.shape(attn_mask)
513
+ if len(mask_shape) == 4 and mask_shape[-1] == self.attention_heads:
514
+ attn_mask = ops.transpose(attn_mask, (0, 3, 1, 2))
515
+ elif len(mask_shape) == 3:
516
+ attn_mask = ops.reshape(attn_mask, (bsz, self.attention_heads, seq_len, seq_len))
517
+
518
+ input_attn_mask = attn_mask
519
+ curr_attn_mask = attn_mask
520
+
521
+ for enc_layer in self.layers_list:
522
+ x, curr_attn_mask, _ = enc_layer(
523
+ x,
524
+ attn_bias=curr_attn_mask,
525
+ padding_mask=padding_mask,
526
+ return_attn=True,
527
+ training=training,
528
+ )
529
+
530
+ if self.final_layer_norm is not None:
531
+ x = self.final_layer_norm(x)
532
+
533
+ # Delta pair representation
534
+ delta_pair_repr = curr_attn_mask - input_attn_mask
535
+
536
+ # Reshape to [bsz, seq_len, seq_len, attention_heads]
537
+ pair_shape = ops.shape(curr_attn_mask)
538
+ if len(pair_shape) == 3:
539
+ curr_attn_mask = ops.reshape(curr_attn_mask, (bsz, self.attention_heads, seq_len, seq_len))
540
+ delta_pair_repr = ops.reshape(delta_pair_repr, (bsz, self.attention_heads, seq_len, seq_len))
541
+
542
+ curr_attn_mask = ops.transpose(curr_attn_mask, (0, 2, 3, 1))
543
+ delta_pair_repr = ops.transpose(delta_pair_repr, (0, 2, 3, 1))
544
+
545
+ if self.final_head_layer_norm is not None:
546
+ delta_pair_repr = self.final_head_layer_norm(delta_pair_repr)
547
+
548
+ return x, curr_attn_mask, delta_pair_repr
549
+
550
+
551
+ # ==============================================================================
552
+ # Uni-Mol Model
553
+ # ==============================================================================
554
+
555
+ class UniMolModel(keras.Model):
556
+ r"""Multi-backend Uni-Mol model for 3D molecular representation learning and property prediction.
557
+
558
+ Supports molecular pretraining, candidate pocket pretraining, crystal, and OLED configurations.
559
+
560
+ Args:
561
+ output_dim (int, optional): Number of task output dimensions / classes. (default: ``2``)
562
+ data_type (str, optional): Data domain (``"molecule"``, ``"protein"``, ``"crystal"``, ``"oled"``). (default: ``"molecule"``)
563
+ vocab_size (int, optional): Vocabulary size for token dictionary. (default: ``512``)
564
+ encoder_layers (int, optional): Number of transformer encoder layers. (default: ``15``)
565
+ encoder_embed_dim (int, optional): Node embedding dimension. (default: ``512``)
566
+ encoder_ffn_embed_dim (int, optional): FFN hidden dimension. (default: ``2048``)
567
+ encoder_attention_heads (int, optional): Number of attention heads. (default: ``64``)
568
+ kernel (str, optional): GBF kernel type (``"gaussian"`` or ``"numerical"``). (default: ``"gaussian"``)
569
+ num_kernel (int, optional): Number of radial kernels. (default: ``128``)
570
+ pooler_dropout (float, optional): Dropout for classification head. (default: ``0.0``)
571
+ activation_fn (str, optional): Activation function name. (default: ``"gelu"``)
572
+ post_ln (bool, optional): Post-LN flag. (default: ``False``)
573
+ **kwargs: Additional model arguments.
574
+
575
+ Example:
576
+ ```python
577
+ import numpy as np
578
+ from k3_node.models import UniMolModel
579
+
580
+ tokens = np.random.randint(1, 64, size=(2, 6)) # atom tokens of 2 molecules with 6 atoms
581
+ coords = np.random.rand(2, 6, 3).astype("float32") * 3.0 # 3D conformations
582
+
583
+ model = UniMolModel(output_dim=2, vocab_size=64, encoder_layers=2, encoder_embed_dim=32, encoder_ffn_embed_dim=64,
584
+ encoder_attention_heads=4, num_kernel=16)
585
+ logits = model(tokens, src_coord=coords) # molecule-level predictions
586
+ print(tuple(logits.shape)) # (2, 2)
587
+ reprs = model(tokens, src_coord=coords, return_repr=True)
588
+ print(tuple(reprs["cls_repr"].shape), tuple(reprs["encoder_rep"].shape)) # (2, 32) (2, 6, 32): molecule and atom embeddings
589
+ ```
590
+ """
591
+
592
+ def __init__(
593
+ self,
594
+ output_dim: int = 2,
595
+ data_type: str = "molecule",
596
+ vocab_size: int = 512,
597
+ encoder_layers: int = 15,
598
+ encoder_embed_dim: int = 512,
599
+ encoder_ffn_embed_dim: int = 2048,
600
+ encoder_attention_heads: int = 64,
601
+ kernel: str = "gaussian",
602
+ num_kernel: int = 128,
603
+ pooler_dropout: float = 0.0,
604
+ activation_fn: str = "gelu",
605
+ post_ln: bool = False,
606
+ **kwargs,
607
+ ):
608
+ super().__init__(**kwargs)
609
+ self.output_dim = output_dim
610
+ self.data_type = data_type
611
+ self.vocab_size = vocab_size
612
+ self.encoder_layers = encoder_layers
613
+ self.encoder_embed_dim = encoder_embed_dim
614
+ self.encoder_ffn_embed_dim = encoder_ffn_embed_dim
615
+ self.encoder_attention_heads = encoder_attention_heads
616
+ self.kernel_type = kernel
617
+ self.num_kernel = num_kernel
618
+ self.pooler_dropout = pooler_dropout
619
+ self.activation_fn_name = activation_fn
620
+ self.post_ln = post_ln
621
+
622
+ self.padding_idx = 0
623
+
624
+ self.embed_tokens = layers.Embedding(vocab_size, encoder_embed_dim, name="embed_tokens")
625
+
626
+ n_edge_type = 1024
627
+ if kernel == "gaussian":
628
+ self.gbf = GaussianLayer(num_kernel, n_edge_type, name="gbf")
629
+ else:
630
+ self.gbf = NumericalEmbed(num_kernel, n_edge_type, activation_fn=activation_fn, name="gbf")
631
+
632
+ self.gbf_proj = NonLinearHead(
633
+ num_kernel, encoder_attention_heads, activation_fn=activation_fn, name="gbf_proj"
634
+ )
635
+
636
+ self.encoder = UniMolTransformerEncoder(
637
+ encoder_layers=encoder_layers,
638
+ embed_dim=encoder_embed_dim,
639
+ ffn_embed_dim=encoder_ffn_embed_dim,
640
+ attention_heads=encoder_attention_heads,
641
+ activation_fn=activation_fn,
642
+ post_ln=post_ln,
643
+ name="encoder",
644
+ )
645
+
646
+ self.classification_head = LinearHead(
647
+ input_dim=encoder_embed_dim,
648
+ num_classes=output_dim,
649
+ pooler_dropout=pooler_dropout,
650
+ name="classification_head",
651
+ )
652
+
653
+ self.pair2coord_proj = NonLinearHead(
654
+ encoder_attention_heads, 1, activation_fn=activation_fn, name="pair2coord_proj"
655
+ )
656
+ self.dist_head = DistanceHead(
657
+ encoder_attention_heads, activation_fn=activation_fn, name="dist_head"
658
+ )
659
+ self.lm_head = MaskLMHead(
660
+ encoder_embed_dim, vocab_size, activation_fn=activation_fn, name="lm_head"
661
+ )
662
+
663
+ def build(self, input_shape=None):
664
+ if not self.built:
665
+ self.embed_tokens.build(None)
666
+ self.gbf.build(None)
667
+ self.gbf_proj.build((None, None, None, self.num_kernel))
668
+ self.encoder.build(None)
669
+ self.classification_head.build((None, self.encoder_embed_dim))
670
+ self.pair2coord_proj.build((None, None, None, self.encoder_attention_heads))
671
+ self.dist_head.build((None, None, None, self.encoder_attention_heads))
672
+ self.lm_head.build((None, None, self.encoder_embed_dim))
673
+ super().build(input_shape)
674
+
675
+ def call(
676
+ self,
677
+ src_tokens,
678
+ src_distance=None,
679
+ src_coord=None,
680
+ src_edge_type=None,
681
+ padding_mask=None,
682
+ return_repr: bool = False,
683
+ return_atomic_reprs: bool = False,
684
+ features_only: bool = False,
685
+ training: bool = False,
686
+ ):
687
+ if isinstance(src_tokens, dict):
688
+ src_distance = src_tokens.get("src_distance", src_distance)
689
+ src_coord = src_tokens.get("src_coord", src_coord)
690
+ src_edge_type = src_tokens.get("src_edge_type", src_edge_type)
691
+ padding_mask = src_tokens.get("padding_mask", padding_mask)
692
+ src_tokens = src_tokens.get("src_tokens", src_tokens.get("tokens"))
693
+ elif isinstance(src_tokens, (tuple, list)):
694
+ if len(src_tokens) == 2:
695
+ src_tokens, src_coord = src_tokens
696
+ elif len(src_tokens) >= 3:
697
+ src_tokens, src_distance, src_coord = src_tokens[:3]
698
+ r"""Forward pass for Uni-Mol.
699
+
700
+ Args:
701
+ src_tokens (Tensor): Token indices of shape ``[batch_size, seq_len]``.
702
+ src_distance (Tensor, optional): Pairwise distances ``[batch_size, seq_len, seq_len]``.
703
+ src_coord (Tensor, optional): 3D coordinates ``[batch_size, seq_len, 3]``.
704
+ src_edge_type (Tensor, optional): Pairwise edge types ``[batch_size, seq_len, seq_len]``.
705
+ padding_mask (Tensor, optional): Padding boolean mask ``[batch_size, seq_len]``.
706
+ return_repr (bool, optional): Return CLS/atomic representations.
707
+ features_only (bool, optional): Only return representations, skipping heads.
708
+ training (bool, optional): Training mode flag.
709
+ """
710
+ shape = ops.shape(src_tokens)
711
+ bsz = shape[0]
712
+ seq_len = shape[1]
713
+
714
+ # Auto-compute padding_mask if not provided
715
+ if padding_mask is None:
716
+ padding_mask = ops.equal(src_tokens, self.padding_idx)
717
+
718
+ # Auto-compute src_distance from src_coord if distance is missing
719
+ if src_distance is None and src_coord is not None:
720
+ diff = ops.expand_dims(src_coord, axis=2) - ops.expand_dims(src_coord, axis=1)
721
+ src_distance = ops.sqrt(ops.sum(ops.power(diff, 2), axis=-1) + 1e-10)
722
+ elif src_distance is None:
723
+ src_distance = ops.zeros((bsz, seq_len, seq_len), dtype="float32")
724
+
725
+ # Auto-compute src_edge_type from tokens if missing
726
+ if src_edge_type is None:
727
+ n_types = 32
728
+ src_edge_type = ops.expand_dims(src_tokens, axis=-1) * n_types + ops.expand_dims(src_tokens, axis=1)
729
+ src_edge_type = ops.cast(ops.mod(src_edge_type, 1024), "int32")
730
+
731
+ # 1. Embeddings & GBF
732
+ x = self.embed_tokens(src_tokens)
733
+ gbf_feature = self.gbf(src_distance, src_edge_type)
734
+ graph_attn_bias = self.gbf_proj(gbf_feature) # [B, N, N, H]
735
+
736
+ # 2. Encoder
737
+ encoder_rep, encoder_pair_rep, delta_pair_rep = self.encoder(
738
+ x,
739
+ attn_mask=graph_attn_bias,
740
+ padding_mask=padding_mask,
741
+ training=training,
742
+ )
743
+
744
+ cls_repr = encoder_rep[:, 0, :]
745
+
746
+ if return_repr:
747
+ res = {"cls_repr": cls_repr, "encoder_rep": encoder_rep, "pair_rep": encoder_pair_rep}
748
+ return res
749
+
750
+ if features_only:
751
+ return encoder_rep, encoder_pair_rep
752
+
753
+ # Classification / Regression logits
754
+ logits = self.classification_head(cls_repr, training=training)
755
+ return logits
756
+
757
+
758
+ class UniMolConfGenModel(UniMolModel):
759
+ r"""Uni-Mol Conformation Generation Model for iterative 3D geometry prediction.
760
+
761
+ Example:
762
+ ```python
763
+ import numpy as np
764
+ from k3_node.models import UniMolConfGenModel
765
+
766
+ tokens = np.random.randint(1, 64, size=(2, 6)) # atom tokens of 2 molecules with 6 atoms
767
+ coords = np.random.rand(2, 6, 3).astype("float32") * 3.0 # 3D conformations
768
+
769
+ model = UniMolConfGenModel(vocab_size=64, encoder_layers=2, encoder_embed_dim=32, encoder_ffn_embed_dim=64,
770
+ encoder_attention_heads=4, num_kernel=16)
771
+ new_coords, pred_dist = model(tokens, src_coord=coords) # refined conformation, predicted distances
772
+ print(tuple(new_coords.shape), tuple(pred_dist.shape)) # (2, 6, 3) (2, 6, 6)
773
+ ```
774
+ """
775
+
776
+ def call(
777
+ self,
778
+ src_tokens,
779
+ src_distance=None,
780
+ src_coord=None,
781
+ src_edge_type=None,
782
+ padding_mask=None,
783
+ training: bool = False,
784
+ ):
785
+ if isinstance(src_tokens, dict):
786
+ src_distance = src_tokens.get("src_distance", src_distance)
787
+ src_coord = src_tokens.get("src_coord", src_coord)
788
+ src_edge_type = src_tokens.get("src_edge_type", src_edge_type)
789
+ padding_mask = src_tokens.get("padding_mask", padding_mask)
790
+ src_tokens = src_tokens.get("src_tokens", src_tokens.get("tokens"))
791
+ elif isinstance(src_tokens, (tuple, list)):
792
+ if len(src_tokens) == 2:
793
+ src_tokens, src_coord = src_tokens
794
+ elif len(src_tokens) >= 3:
795
+ src_tokens, src_distance, src_coord = src_tokens[:3]
796
+
797
+ if padding_mask is None:
798
+ padding_mask = ops.equal(src_tokens, self.padding_idx)
799
+
800
+ shape = ops.shape(src_tokens)
801
+ bsz = shape[0]
802
+ seq_len = shape[1]
803
+
804
+ if src_coord is None:
805
+ src_coord = ops.zeros((bsz, seq_len, 3), dtype="float32")
806
+
807
+ if src_distance is None:
808
+ diff = ops.expand_dims(src_coord, axis=2) - ops.expand_dims(src_coord, axis=1)
809
+ src_distance = ops.sqrt(ops.sum(ops.power(diff, 2), axis=-1) + 1e-10)
810
+
811
+ if src_edge_type is None:
812
+ src_edge_type = ops.cast(ops.mod(ops.expand_dims(src_tokens, axis=-1) * 32 + ops.expand_dims(src_tokens, axis=1), 1024), "int32")
813
+
814
+ x = self.embed_tokens(src_tokens)
815
+ gbf_feature = self.gbf(src_distance, src_edge_type)
816
+ graph_attn_bias = self.gbf_proj(gbf_feature)
817
+
818
+ encoder_rep, encoder_pair_rep, delta_pair_rep = self.encoder(
819
+ x,
820
+ attn_mask=graph_attn_bias,
821
+ padding_mask=padding_mask,
822
+ training=training,
823
+ )
824
+
825
+ # Coordinate update from delta_pair_rep
826
+ attn_probs = self.pair2coord_proj(delta_pair_rep) # [B, N, N, 1]
827
+ delta_pos = ops.expand_dims(src_coord, axis=1) - ops.expand_dims(src_coord, axis=2)
828
+ coord_update = delta_pos * attn_probs
829
+ updated_coord = src_coord + ops.sum(coord_update, axis=2)
830
+
831
+ pred_dist = self.dist_head(encoder_pair_rep)
832
+ return updated_coord, pred_dist
833
+
834
+
835
+ class UniMolDockingModel(UniMolModel):
836
+ r"""Uni-Mol Protein-Ligand Binding Pose Prediction Model.
837
+
838
+ Example:
839
+ ```python
840
+ import numpy as np
841
+ from k3_node.models import UniMolDockingModel
842
+
843
+ tokens = np.random.randint(1, 64, size=(2, 6)) # atom tokens of 2 molecules with 6 atoms
844
+ coords = np.random.rand(2, 6, 3).astype("float32") * 3.0 # 3D conformations
845
+
846
+ model = UniMolDockingModel(vocab_size=64, encoder_layers=2, encoder_embed_dim=32, encoder_ffn_embed_dim=64,
847
+ encoder_attention_heads=4, num_kernel=16)
848
+ pose, pred_dist = model(tokens, src_coord=coords)
849
+ print(tuple(pose.shape), tuple(pred_dist.shape)) # (2, 6, 3) (2, 6, 6)
850
+ ```
851
+ """
852
+
853
+ def __init__(self, **kwargs):
854
+ super().__init__(**kwargs)
855
+ # Cross-layer coordinate prediction head
856
+ self.docking_coord_head = NonLinearHead(
857
+ self.encoder_attention_heads, 1, activation_fn="gelu", name="docking_coord_head"
858
+ )
859
+
860
+ def build(self, input_shape=None):
861
+ if not self.built:
862
+ self.docking_coord_head.build((None, None, None, self.encoder_attention_heads))
863
+ super().build(input_shape)
864
+
865
+ def call(
866
+ self,
867
+ src_tokens,
868
+ src_distance=None,
869
+ src_coord=None,
870
+ src_edge_type=None,
871
+ padding_mask=None,
872
+ training: bool = False,
873
+ ):
874
+ if isinstance(src_tokens, dict):
875
+ src_distance = src_tokens.get("src_distance", src_distance)
876
+ src_coord = src_tokens.get("src_coord", src_coord)
877
+ src_edge_type = src_tokens.get("src_edge_type", src_edge_type)
878
+ padding_mask = src_tokens.get("padding_mask", padding_mask)
879
+ src_tokens = src_tokens.get("src_tokens", src_tokens.get("tokens"))
880
+ elif isinstance(src_tokens, (tuple, list)):
881
+ if len(src_tokens) == 2:
882
+ src_tokens, src_coord = src_tokens
883
+ elif len(src_tokens) >= 3:
884
+ src_tokens, src_distance, src_coord = src_tokens[:3]
885
+
886
+ if padding_mask is None:
887
+ padding_mask = ops.equal(src_tokens, self.padding_idx)
888
+
889
+ shape = ops.shape(src_tokens)
890
+ bsz = shape[0]
891
+ seq_len = shape[1]
892
+
893
+ if src_coord is None:
894
+ src_coord = ops.zeros((bsz, seq_len, 3), dtype="float32")
895
+
896
+ if src_distance is None:
897
+ diff = ops.expand_dims(src_coord, axis=2) - ops.expand_dims(src_coord, axis=1)
898
+ src_distance = ops.sqrt(ops.sum(ops.power(diff, 2), axis=-1) + 1e-10)
899
+
900
+ if src_edge_type is None:
901
+ src_edge_type = ops.cast(ops.mod(ops.expand_dims(src_tokens, axis=-1) * 32 + ops.expand_dims(src_tokens, axis=1), 1024), "int32")
902
+
903
+ x = self.embed_tokens(src_tokens)
904
+ gbf_feature = self.gbf(src_distance, src_edge_type)
905
+ graph_attn_bias = self.gbf_proj(gbf_feature)
906
+
907
+ encoder_rep, encoder_pair_rep, delta_pair_rep = self.encoder(
908
+ x,
909
+ attn_mask=graph_attn_bias,
910
+ padding_mask=padding_mask,
911
+ training=training,
912
+ )
913
+
914
+ probs = self.docking_coord_head(delta_pair_rep)
915
+ diff_coord = ops.expand_dims(src_coord, axis=1) - ops.expand_dims(src_coord, axis=2)
916
+ pose_update = src_coord + ops.sum(diff_coord * probs, axis=2)
917
+ pred_dist = self.dist_head(encoder_pair_rep)
918
+ return pose_update, pred_dist
919
+
920
+
921
+ # ==============================================================================
922
+ # Helper Functions: Download & Load Weights
923
+ # ==============================================================================
924
+
925
+ def download_unimol_checkpoint(
926
+ name: str = "mol_pre_no_h",
927
+ folder: str = "checkpoints",
928
+ log: bool = True,
929
+ ) -> str:
930
+ r"""Downloads a pre-trained Uni-Mol checkpoint (.pt).
931
+
932
+ Args:
933
+ name (str): Pretrained checkpoint name or alias (e.g. ``"mol_pre_no_h"``,
934
+ ``"mol_pre_all_h"``, ``"pocket_pre"``, ``"mp_all_h"``, ``"oled_pre_no_h"``,
935
+ ``"qm9"``, ``"drugs"``, ``"binding_pose"``).
936
+ folder (str, optional): Target directory to save the checkpoint. (default: ``"checkpoints"``)
937
+ log (bool, optional): Whether to print download progress. (default: ``True``)
938
+
939
+ Returns:
940
+ str: Absolute path to the downloaded checkpoint file.
941
+ """
942
+ clean_name = name.strip().lower()
943
+ if clean_name in UNIMOL_ALIASES:
944
+ clean_name = UNIMOL_ALIASES[clean_name]
945
+
946
+ if clean_name not in UNIMOL_PRETRAINED_URLS:
947
+ for k, v in UNIMOL_PRETRAINED_URLS.items():
948
+ if v["filename"] == name or k.lower() == clean_name:
949
+ clean_name = k
950
+ break
951
+
952
+ if clean_name not in UNIMOL_PRETRAINED_URLS:
953
+ raise ValueError(
954
+ f"Unknown Uni-Mol checkpoint '{name}'. Available: {list(UNIMOL_PRETRAINED_URLS.keys())}"
955
+ )
956
+
957
+ info = UNIMOL_PRETRAINED_URLS[clean_name]
958
+ filename = info["filename"]
959
+ local_path = os.path.join(folder, filename)
960
+
961
+ if os.path.exists(local_path):
962
+ return local_path
963
+
964
+ # Check alternative local paths
965
+ alt_paths = [
966
+ os.path.join("Uni-Mol", "unimol", filename),
967
+ os.path.join("Uni-Mol", "unimol_tools", "unimol_tools", "weights", filename),
968
+ ]
969
+ for alt in alt_paths:
970
+ if os.path.exists(alt):
971
+ return alt
972
+
973
+ return download_url(info["url"], folder=folder, filename=filename, log=log)
974
+
975
+
976
+ def load_unimol_weights(
977
+ model: UniMolModel,
978
+ checkpoint_path: Optional[str] = None,
979
+ pretrained_name: Optional[str] = None,
980
+ folder: str = "checkpoints",
981
+ download: bool = True,
982
+ ) -> UniMolModel:
983
+ r"""Loads pre-trained weights from a PyTorch checkpoint (.pt) into a Keras 3 UniMolModel.
984
+
985
+ Args:
986
+ model (UniMolModel): The target UniMolModel instance.
987
+ checkpoint_path (str, optional): Local path to .pt file or checkpoint name.
988
+ pretrained_name (str, optional): Pretrained model identifier.
989
+ folder (str, optional): Directory to store downloaded checkpoints. (default: ``"checkpoints"``)
990
+ download (bool, optional): Whether to download checkpoint if missing locally. (default: ``True``)
991
+
992
+ Returns:
993
+ UniMolModel: The model with loaded weights.
994
+ """
995
+ path_to_load = checkpoint_path
996
+
997
+ candidate = pretrained_name or checkpoint_path
998
+ if candidate:
999
+ clean = candidate.strip().lower()
1000
+ if clean in UNIMOL_ALIASES:
1001
+ candidate = UNIMOL_ALIASES[clean]
1002
+
1003
+ if path_to_load is None:
1004
+ if candidate is None:
1005
+ raise ValueError("Either checkpoint_path or pretrained_name must be specified.")
1006
+ if candidate in UNIMOL_PRETRAINED_URLS:
1007
+ fname = UNIMOL_PRETRAINED_URLS[candidate]["filename"]
1008
+ local_target = os.path.join(folder, fname)
1009
+ if os.path.isfile(local_target):
1010
+ path_to_load = local_target
1011
+ elif download:
1012
+ path_to_load = download_unimol_checkpoint(candidate, folder=folder)
1013
+ else:
1014
+ raise FileNotFoundError(f"Checkpoint for '{candidate}' not found at '{local_target}'.")
1015
+ else:
1016
+ raise ValueError(f"Unknown checkpoint '{candidate}'.")
1017
+ elif not os.path.isfile(path_to_load) and download and candidate in UNIMOL_PRETRAINED_URLS:
1018
+ path_to_load = download_unimol_checkpoint(candidate, folder=folder)
1019
+
1020
+ import torch
1021
+
1022
+ state = torch.load(path_to_load, map_location="cpu")
1023
+ if isinstance(state, dict):
1024
+ if "model" in state:
1025
+ state_dict = state["model"]
1026
+ elif "model_state_dict" in state:
1027
+ state_dict = state["model_state_dict"]
1028
+ else:
1029
+ state_dict = state
1030
+ else:
1031
+ raise ValueError(f"Expected dict/state_dict in checkpoint, got {type(state)}")
1032
+
1033
+ if not model.built:
1034
+ model.build(None)
1035
+
1036
+ def _to_tensor(t):
1037
+ if hasattr(t, "detach"):
1038
+ t = t.detach()
1039
+ if hasattr(t, "cpu"):
1040
+ t = t.cpu()
1041
+ if hasattr(t, "numpy"):
1042
+ t = t.numpy()
1043
+ return ops.convert_to_tensor(np.array(t, dtype=np.float32), dtype="float32")
1044
+
1045
+ # 1. Embeddings
1046
+ if "embed_tokens.weight" in state_dict:
1047
+ w = state_dict["embed_tokens.weight"]
1048
+ model.embed_tokens.embeddings.assign(_to_tensor(w))
1049
+
1050
+ # 2. GBF
1051
+ if "gbf.means.weight" in state_dict:
1052
+ model.gbf.means.embeddings.assign(_to_tensor(state_dict["gbf.means.weight"]))
1053
+ if "gbf.stds.weight" in state_dict:
1054
+ model.gbf.stds.embeddings.assign(_to_tensor(state_dict["gbf.stds.weight"]))
1055
+ if "gbf.mul.weight" in state_dict:
1056
+ model.gbf.mul.embeddings.assign(_to_tensor(state_dict["gbf.mul.weight"]))
1057
+ if "gbf.bias.weight" in state_dict:
1058
+ model.gbf.bias.embeddings.assign(_to_tensor(state_dict["gbf.bias.weight"]))
1059
+
1060
+ # 3. GBF Proj
1061
+ if "gbf_proj.linear1.weight" in state_dict:
1062
+ model.gbf_proj.linear1.kernel.assign(_to_tensor(state_dict["gbf_proj.linear1.weight"].t()))
1063
+ if "gbf_proj.linear1.bias" in state_dict:
1064
+ model.gbf_proj.linear1.bias.assign(_to_tensor(state_dict["gbf_proj.linear1.bias"]))
1065
+ if "gbf_proj.linear2.weight" in state_dict:
1066
+ model.gbf_proj.linear2.kernel.assign(_to_tensor(state_dict["gbf_proj.linear2.weight"].t()))
1067
+ if "gbf_proj.linear2.bias" in state_dict:
1068
+ model.gbf_proj.linear2.bias.assign(_to_tensor(state_dict["gbf_proj.linear2.bias"]))
1069
+
1070
+ # 4. Encoder Layers
1071
+ for i, enc_layer in enumerate(model.encoder.layers_list):
1072
+ prefix = f"encoder.layers.{i}"
1073
+ alt_prefix = f"layers.{i}"
1074
+
1075
+ def _get(name):
1076
+ if f"{prefix}.{name}" in state_dict:
1077
+ return state_dict[f"{prefix}.{name}"]
1078
+ elif f"{alt_prefix}.{name}" in state_dict:
1079
+ return state_dict[f"{alt_prefix}.{name}"]
1080
+ return None
1081
+
1082
+ # Self-Attention
1083
+ in_w = _get("self_attn.in_proj.weight")
1084
+ if in_w is not None:
1085
+ enc_layer.self_attn.in_proj.kernel.assign(_to_tensor(in_w.t()))
1086
+ in_b = _get("self_attn.in_proj.bias")
1087
+ if in_b is not None:
1088
+ enc_layer.self_attn.in_proj.bias.assign(_to_tensor(in_b))
1089
+
1090
+ out_w = _get("self_attn.out_proj.weight")
1091
+ if out_w is not None:
1092
+ enc_layer.self_attn.out_proj.kernel.assign(_to_tensor(out_w.t()))
1093
+ out_b = _get("self_attn.out_proj.bias")
1094
+ if out_b is not None:
1095
+ enc_layer.self_attn.out_proj.bias.assign(_to_tensor(out_b))
1096
+
1097
+ # Layer norms
1098
+ attn_ln_w = _get("self_attn_layer_norm.weight")
1099
+ if attn_ln_w is not None and enc_layer.self_attn_layer_norm.gamma is not None:
1100
+ enc_layer.self_attn_layer_norm.gamma.assign(_to_tensor(attn_ln_w))
1101
+ attn_ln_b = _get("self_attn_layer_norm.bias")
1102
+ if attn_ln_b is not None and enc_layer.self_attn_layer_norm.beta is not None:
1103
+ enc_layer.self_attn_layer_norm.beta.assign(_to_tensor(attn_ln_b))
1104
+
1105
+ # FFN
1106
+ fc1_w = _get("fc1.weight")
1107
+ if fc1_w is not None:
1108
+ enc_layer.fc1.kernel.assign(_to_tensor(fc1_w.t()))
1109
+ fc1_b = _get("fc1.bias")
1110
+ if fc1_b is not None:
1111
+ enc_layer.fc1.bias.assign(_to_tensor(fc1_b))
1112
+
1113
+ fc2_w = _get("fc2.weight")
1114
+ if fc2_w is not None:
1115
+ enc_layer.fc2.kernel.assign(_to_tensor(fc2_w.t()))
1116
+ fc2_b = _get("fc2.bias")
1117
+ if fc2_b is not None:
1118
+ enc_layer.fc2.bias.assign(_to_tensor(fc2_b))
1119
+
1120
+ final_ln_w = _get("final_layer_norm.weight")
1121
+ if final_ln_w is not None and enc_layer.final_layer_norm.gamma is not None:
1122
+ enc_layer.final_layer_norm.gamma.assign(_to_tensor(final_ln_w))
1123
+ final_ln_b = _get("final_layer_norm.bias")
1124
+ if final_ln_b is not None and enc_layer.final_layer_norm.beta is not None:
1125
+ enc_layer.final_layer_norm.beta.assign(_to_tensor(final_ln_b))
1126
+
1127
+ # Encoder global norms
1128
+ for key, attr in [
1129
+ ("encoder.emb_layer_norm.weight", model.encoder.emb_layer_norm.gamma),
1130
+ ("encoder.emb_layer_norm.bias", model.encoder.emb_layer_norm.beta),
1131
+ ]:
1132
+ if key in state_dict and attr is not None:
1133
+ attr.assign(_to_tensor(state_dict[key]))
1134
+
1135
+ if model.encoder.final_layer_norm is not None:
1136
+ if "encoder.final_layer_norm.weight" in state_dict:
1137
+ model.encoder.final_layer_norm.gamma.assign(_to_tensor(state_dict["encoder.final_layer_norm.weight"]))
1138
+ if "encoder.final_layer_norm.bias" in state_dict:
1139
+ model.encoder.final_layer_norm.beta.assign(_to_tensor(state_dict["encoder.final_layer_norm.bias"]))
1140
+
1141
+ if model.encoder.final_head_layer_norm is not None:
1142
+ if "encoder.final_head_layer_norm.weight" in state_dict:
1143
+ model.encoder.final_head_layer_norm.gamma.assign(_to_tensor(state_dict["encoder.final_head_layer_norm.weight"]))
1144
+ if "encoder.final_head_layer_norm.bias" in state_dict:
1145
+ model.encoder.final_head_layer_norm.beta.assign(_to_tensor(state_dict["encoder.final_head_layer_norm.bias"]))
1146
+
1147
+ # Classification / linear head
1148
+ if hasattr(model, "classification_head"):
1149
+ for k in ["classification_head.out_proj.weight", "classification_heads.target.out_proj.weight"]:
1150
+ if k in state_dict and hasattr(model.classification_head, "out_proj"):
1151
+ model.classification_head.out_proj.kernel.assign(_to_tensor(state_dict[k].t()))
1152
+ for k in ["classification_head.out_proj.bias", "classification_heads.target.out_proj.bias"]:
1153
+ if k in state_dict and hasattr(model.classification_head, "out_proj"):
1154
+ model.classification_head.out_proj.bias.assign(_to_tensor(state_dict[k]))
1155
+
1156
+ return model