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,85 @@
1
+ import json
2
+ import warnings
3
+ from itertools import chain
4
+ from typing import Callable, List, Optional
5
+ import numpy as np
6
+ from keras import ops
7
+
8
+ from k3_node.data import Data, InMemoryDataset
9
+ from k3_node.io import fs
10
+ from k3_node.transforms.utils import to_undirected
11
+
12
+
13
+ class WikiCS(InMemoryDataset):
14
+ r"""The semi-supervised Wikipedia-based dataset from the
15
+ "Wiki-CS: A Wikipedia-Based Benchmark for Graph Neural Networks" paper.
16
+
17
+ Args:
18
+ root (str): Root directory where the dataset should be saved.
19
+ transform (callable, optional): Transform function.
20
+ pre_transform (callable, optional): Pre-transform function.
21
+ is_undirected (bool, optional): Whether the graph is undirected. (default: True)
22
+ force_reload (bool, optional): Whether to re-process the dataset.
23
+ """
24
+
25
+ url = "https://github.com/pmernyei/wiki-cs-dataset/raw/master/dataset"
26
+
27
+ def __init__(
28
+ self,
29
+ root: str,
30
+ transform: Optional[Callable] = None,
31
+ pre_transform: Optional[Callable] = None,
32
+ is_undirected: Optional[bool] = None,
33
+ force_reload: bool = False,
34
+ ):
35
+ if is_undirected is None:
36
+ is_undirected = True
37
+ self.is_undirected = is_undirected
38
+ super().__init__(root, transform, pre_transform, force_reload=force_reload)
39
+ self.load(self.processed_paths[0])
40
+
41
+ @property
42
+ def raw_file_names(self) -> List[str]:
43
+ return ["data.json"]
44
+
45
+ @property
46
+ def processed_file_names(self) -> str:
47
+ return "data_undirected.pt" if self.is_undirected else "data.pt"
48
+
49
+ def download(self):
50
+ for name in self.raw_file_names:
51
+ fs.cp(f"{self.url}/{name}", self.raw_dir)
52
+
53
+ def process(self):
54
+ with open(self.raw_paths[0]) as f:
55
+ data = json.load(f)
56
+
57
+ x = np.array(data["features"], dtype=np.float32)
58
+ y = np.array(data["labels"], dtype=np.int64)
59
+
60
+ edges = [[(i, j) for j in js] for i, js in enumerate(data["links"])]
61
+ edges = list(chain(*edges))
62
+ edge_index = np.array(edges, dtype=np.int64).T
63
+ if self.is_undirected:
64
+ edge_index = to_undirected(edge_index, num_nodes=x.shape[0])
65
+
66
+ train_mask = np.array(data["train_masks"], dtype=bool).T
67
+ val_mask = np.array(data["val_masks"], dtype=bool).T
68
+ test_mask = np.array(data["test_mask"], dtype=bool)
69
+ stopping_mask = np.array(data["stopping_masks"], dtype=bool).T
70
+
71
+ data_obj = Data(
72
+ x=ops.convert_to_tensor(x, dtype="float32"),
73
+ y=ops.convert_to_tensor(y, dtype="int64"),
74
+ edge_index=ops.convert_to_tensor(edge_index, dtype="int64"),
75
+ train_mask=ops.convert_to_tensor(train_mask, dtype="bool"),
76
+ val_mask=ops.convert_to_tensor(val_mask, dtype="bool"),
77
+ test_mask=ops.convert_to_tensor(test_mask, dtype="bool"),
78
+ stopping_mask=ops.convert_to_tensor(stopping_mask, dtype="bool"),
79
+ )
80
+
81
+ if self.pre_transform is not None:
82
+ data_obj = self.pre_transform(data_obj)
83
+
84
+ self.save([data_obj], self.processed_paths[0])
85
+
@@ -0,0 +1,184 @@
1
+ from typing import Callable, List, Optional
2
+ import numpy as np
3
+ from keras import ops
4
+
5
+ from k3_node.data import Data, InMemoryDataset
6
+ from k3_node.io import fs
7
+
8
+
9
+ class WordNet18(InMemoryDataset):
10
+ r"""The WordNet18 dataset containing 40,943 entities, 18 relations and 151,442 fact triplets."""
11
+
12
+ url = "https://raw.githubusercontent.com/villmow/datasets_knowledge_embedding/master/WN18/original"
13
+
14
+ def __init__(
15
+ self,
16
+ root: str,
17
+ transform: Optional[Callable] = None,
18
+ pre_transform: Optional[Callable] = None,
19
+ force_reload: bool = False,
20
+ ):
21
+ super().__init__(root, transform, pre_transform, force_reload=force_reload)
22
+ self.load(self.processed_paths[0])
23
+
24
+ @property
25
+ def raw_file_names(self) -> List[str]:
26
+ return ["train.txt", "valid.txt", "test.txt"]
27
+
28
+ @property
29
+ def processed_file_names(self) -> str:
30
+ return "data.pt"
31
+
32
+ def download(self):
33
+ for filename in self.raw_file_names:
34
+ fs.cp(f"{self.url}/{filename}", self.raw_dir)
35
+
36
+ def process(self):
37
+ srcs, dsts, edge_types = [], [], []
38
+ for path in self.raw_paths:
39
+ with open(path) as f:
40
+ edges = [int(x) for x in f.read().split()[1:]]
41
+ edge = np.array(edges, dtype=np.int64)
42
+ srcs.append(edge[::3])
43
+ dsts.append(edge[1::3])
44
+ edge_types.append(edge[2::3])
45
+
46
+ src = np.concatenate(srcs, axis=0)
47
+ dst = np.concatenate(dsts, axis=0)
48
+ edge_type = np.concatenate(edge_types, axis=0)
49
+
50
+ n_train = len(srcs[0])
51
+ n_val = len(srcs[1])
52
+ n_test = len(srcs[2])
53
+
54
+ train_mask = np.zeros(len(src), dtype=bool)
55
+ train_mask[:n_train] = True
56
+ val_mask = np.zeros(len(src), dtype=bool)
57
+ val_mask[n_train : n_train + n_val] = True
58
+ test_mask = np.zeros(len(src), dtype=bool)
59
+ test_mask[n_train + n_val :] = True
60
+
61
+ num_nodes = int(max(src.max(), dst.max())) + 1
62
+ perm = np.argsort(num_nodes * src + dst)
63
+
64
+ edge_index = np.stack([src[perm], dst[perm]], axis=0)
65
+ edge_type = edge_type[perm]
66
+ train_mask = train_mask[perm]
67
+ val_mask = val_mask[perm]
68
+ test_mask = test_mask[perm]
69
+
70
+ data = Data(
71
+ edge_index=ops.convert_to_tensor(edge_index, dtype="int64"),
72
+ edge_type=ops.convert_to_tensor(edge_type, dtype="int64"),
73
+ train_mask=ops.convert_to_tensor(train_mask, dtype="bool"),
74
+ val_mask=ops.convert_to_tensor(val_mask, dtype="bool"),
75
+ test_mask=ops.convert_to_tensor(test_mask, dtype="bool"),
76
+ num_nodes=num_nodes,
77
+ )
78
+
79
+ if self.pre_transform is not None:
80
+ data = self.pre_transform(data)
81
+
82
+ self.save([data], self.processed_paths[0])
83
+
84
+
85
+ class WordNet18RR(InMemoryDataset):
86
+ r"""The WordNet18RR dataset."""
87
+
88
+ url = "https://raw.githubusercontent.com/villmow/datasets_knowledge_embedding/master/WN18RR/original"
89
+
90
+ edge2id = {
91
+ "_also_see": 0,
92
+ "_derivationally_related_form": 1,
93
+ "_has_part": 2,
94
+ "_hypernym": 3,
95
+ "_instance_hypernym": 4,
96
+ "_member_meronym": 5,
97
+ "_member_of_domain_region": 6,
98
+ "_member_of_domain_usage": 7,
99
+ "_similar_to": 8,
100
+ "_synset_domain_topic_of": 9,
101
+ "_verb_group": 10,
102
+ }
103
+
104
+ def __init__(
105
+ self,
106
+ root: str,
107
+ transform: Optional[Callable] = None,
108
+ pre_transform: Optional[Callable] = None,
109
+ force_reload: bool = False,
110
+ ):
111
+ super().__init__(root, transform, pre_transform, force_reload=force_reload)
112
+ self.load(self.processed_paths[0])
113
+
114
+ @property
115
+ def raw_file_names(self) -> List[str]:
116
+ return ["train.txt", "valid.txt", "test.txt"]
117
+
118
+ @property
119
+ def processed_file_names(self) -> str:
120
+ return "data.pt"
121
+
122
+ def download(self):
123
+ for filename in self.raw_file_names:
124
+ fs.cp(f"{self.url}/{filename}", self.raw_dir)
125
+
126
+ @staticmethod
127
+ def _node_id(node2id, name):
128
+ if name not in node2id:
129
+ node2id[name] = len(node2id)
130
+ return node2id[name]
131
+
132
+ def process(self):
133
+ # Entities are WordNet synset offsets; map them to consecutive ids, as PyG does
134
+ node2id = {}
135
+ srcs, dsts, edge_types = [], [], []
136
+ for path in self.raw_paths:
137
+ with open(path) as f:
138
+ lines = [line.split() for line in f.read().split("\n")[:-1]]
139
+ src, dst = [], []
140
+ for h, _, t in lines:
141
+ src.append(self._node_id(node2id, h))
142
+ dst.append(self._node_id(node2id, t))
143
+ rel = [self.edge2id[r] for _, r, _ in lines]
144
+ srcs.append(np.array(src, dtype=np.int64))
145
+ dsts.append(np.array(dst, dtype=np.int64))
146
+ edge_types.append(np.array(rel, dtype=np.int64))
147
+
148
+ src = np.concatenate(srcs, axis=0)
149
+ dst = np.concatenate(dsts, axis=0)
150
+ edge_type = np.concatenate(edge_types, axis=0)
151
+
152
+ n_train = len(srcs[0])
153
+ n_val = len(srcs[1])
154
+
155
+ train_mask = np.zeros(len(src), dtype=bool)
156
+ train_mask[:n_train] = True
157
+ val_mask = np.zeros(len(src), dtype=bool)
158
+ val_mask[n_train : n_train + n_val] = True
159
+ test_mask = np.zeros(len(src), dtype=bool)
160
+ test_mask[n_train + n_val :] = True
161
+
162
+ num_nodes = len(node2id)
163
+ perm = np.argsort(num_nodes * src + dst)
164
+
165
+ edge_index = np.stack([src[perm], dst[perm]], axis=0)
166
+ edge_type = edge_type[perm]
167
+ train_mask = train_mask[perm]
168
+ val_mask = val_mask[perm]
169
+ test_mask = test_mask[perm]
170
+
171
+ data = Data(
172
+ edge_index=ops.convert_to_tensor(edge_index, dtype="int64"),
173
+ edge_type=ops.convert_to_tensor(edge_type, dtype="int64"),
174
+ train_mask=ops.convert_to_tensor(train_mask, dtype="bool"),
175
+ val_mask=ops.convert_to_tensor(val_mask, dtype="bool"),
176
+ test_mask=ops.convert_to_tensor(test_mask, dtype="bool"),
177
+ num_nodes=num_nodes,
178
+ )
179
+
180
+ if self.pre_transform is not None:
181
+ data = self.pre_transform(data)
182
+
183
+ self.save([data], self.processed_paths[0])
184
+
@@ -0,0 +1,37 @@
1
+ """Tabular-to-Graph ETL (Extract, Transform, Load) pipelines for K3-Node."""
2
+
3
+ from k3_node.etl.encoders import (
4
+ NumericalEncoder,
5
+ CategoricalEncoder,
6
+ TabularEncoder,
7
+ )
8
+ from k3_node.etl.graph_builders import (
9
+ KNNGraphBuilder,
10
+ SimilarityGraphBuilder,
11
+ SharedEntityGraphBuilder,
12
+ SequentialGraphBuilder,
13
+ )
14
+ from k3_node.etl.table_to_graph import (
15
+ TableToGraph,
16
+ TabularToGraph,
17
+ table_to_graph,
18
+ )
19
+ from k3_node.etl.relational_to_graph import (
20
+ RelationalToGraph,
21
+ relational_to_graph,
22
+ )
23
+
24
+ __all__ = [
25
+ "NumericalEncoder",
26
+ "CategoricalEncoder",
27
+ "TabularEncoder",
28
+ "KNNGraphBuilder",
29
+ "SimilarityGraphBuilder",
30
+ "SharedEntityGraphBuilder",
31
+ "SequentialGraphBuilder",
32
+ "TableToGraph",
33
+ "TabularToGraph",
34
+ "table_to_graph",
35
+ "RelationalToGraph",
36
+ "relational_to_graph",
37
+ ]
@@ -0,0 +1,248 @@
1
+ """Feature preprocessors and encoders for tabular data."""
2
+
3
+ from typing import Any, Dict, List, Optional, Sequence, Union
4
+ import numpy as np
5
+
6
+
7
+ class NumericalEncoder:
8
+ r"""Encodes numerical tabular features with scaling and missing value imputation.
9
+
10
+ Args:
11
+ strategy: Scaling method (``"standard"``, ``"minmax"``, ``"log1p"``, or ``"none"``).
12
+ (default: ``"standard"``)
13
+ impute_strategy: How to fill missing values / NaNs (``"mean"``, ``"median"``,
14
+ ``"zero"``, or float constant). (default: ``"mean"``)
15
+ """
16
+
17
+ def __init__(self, strategy: str = "standard", impute_strategy: Union[str, float] = "mean"):
18
+ self.strategy = strategy.lower()
19
+ self.impute_strategy = impute_strategy
20
+ self.mean_: Optional[float] = None
21
+ self.std_: Optional[float] = None
22
+ self.min_: Optional[float] = None
23
+ self.max_: Optional[float] = None
24
+ self.fill_value_: float = 0.0
25
+
26
+ def fit(self, values: Sequence[Any]):
27
+ arr = np.asarray(values, dtype=np.float64).flatten()
28
+ valid = arr[~np.isnan(arr)]
29
+ if len(valid) == 0:
30
+ valid = np.array([0.0])
31
+
32
+ if self.impute_strategy == "mean":
33
+ self.fill_value_ = float(np.mean(valid))
34
+ elif self.impute_strategy == "median":
35
+ self.fill_value_ = float(np.median(valid))
36
+ elif self.impute_strategy == "min":
37
+ self.fill_value_ = float(np.min(valid))
38
+ elif self.impute_strategy == "zero":
39
+ self.fill_value_ = 0.0
40
+ elif isinstance(self.impute_strategy, (int, float)):
41
+ self.fill_value_ = float(self.impute_strategy)
42
+ else:
43
+ self.fill_value_ = float(np.mean(valid))
44
+
45
+ self.mean_ = float(np.mean(valid))
46
+ self.std_ = float(np.std(valid))
47
+ if self.std_ < 1e-8:
48
+ self.std_ = 1.0
49
+
50
+ self.min_ = float(np.min(valid))
51
+ self.max_ = float(np.max(valid))
52
+ if abs(self.max_ - self.min_) < 1e-8:
53
+ self.max_ = self.min_ + 1.0
54
+
55
+ return self
56
+
57
+ def transform(self, values: Sequence[Any]) -> np.ndarray:
58
+ arr = np.asarray(values, dtype=np.float32).flatten()
59
+ arr = np.nan_to_num(arr, nan=self.fill_value_)
60
+
61
+ if self.strategy == "standard":
62
+ mean = self.mean_ if self.mean_ is not None else 0.0
63
+ std = self.std_ if self.std_ is not None else 1.0
64
+ res = (arr - mean) / std
65
+ elif self.strategy == "minmax":
66
+ vmin = self.min_ if self.min_ is not None else 0.0
67
+ vmax = self.max_ if self.max_ is not None else 1.0
68
+ res = np.clip((arr - vmin) / (vmax - vmin), 0.0, 1.0)
69
+ elif self.strategy == "log1p":
70
+ res = np.log1p(np.maximum(arr, 0.0))
71
+ elif self.strategy == "none":
72
+ res = arr
73
+ else:
74
+ raise ValueError(f"Unknown numerical strategy '{self.strategy}'.")
75
+
76
+ return res.reshape(-1, 1)
77
+
78
+ def fit_transform(self, values: Sequence[Any]) -> np.ndarray:
79
+ return self.fit(values).transform(values)
80
+
81
+
82
+ class CategoricalEncoder:
83
+ r"""Encodes categorical strings or integer values into one-hot or ordinal representations.
84
+
85
+ Args:
86
+ strategy: Encoding method (``"onehot"``, ``"ordinal"``, or ``"hash"``).
87
+ (default: ``"onehot"``)
88
+ handle_unknown: How to handle unseen categories during transform (``"ignore"``,
89
+ ``"error"``, or ``"use_encoded_value"``). (default: ``"ignore"``)
90
+ unknown_value: Numerical value assigned to unseen categories when using ordinal encoding.
91
+ (default: ``-1``)
92
+ hash_dim: Output dimension when using ``"hash"`` strategy. (default: ``16``)
93
+ """
94
+
95
+ def __init__(
96
+ self,
97
+ strategy: str = "onehot",
98
+ handle_unknown: str = "ignore",
99
+ unknown_value: int = -1,
100
+ hash_dim: int = 16,
101
+ ):
102
+ self.strategy = strategy.lower()
103
+ self.handle_unknown = handle_unknown.lower()
104
+ self.unknown_value = unknown_value
105
+ self.hash_dim = hash_dim
106
+ self.vocab_: Dict[Any, int] = {}
107
+ self.inv_vocab_: List[Any] = []
108
+
109
+ def fit(self, values: Sequence[Any]):
110
+ arr = [str(v) if v is not None and not (isinstance(v, float) and np.isnan(v)) else "__MISSING__" for v in values]
111
+ unique_cats = sorted(list(set(arr)))
112
+ self.vocab_ = {cat: idx for idx, cat in enumerate(unique_cats)}
113
+ self.inv_vocab_ = unique_cats
114
+ return self
115
+
116
+ def transform(self, values: Sequence[Any]) -> np.ndarray:
117
+ arr = [str(v) if v is not None and not (isinstance(v, float) and np.isnan(v)) else "__MISSING__" for v in values]
118
+ num_samples = len(arr)
119
+
120
+ if self.strategy == "onehot":
121
+ num_classes = len(self.vocab_)
122
+ if num_classes == 0:
123
+ return np.zeros((num_samples, 1), dtype=np.float32)
124
+ out = np.zeros((num_samples, num_classes), dtype=np.float32)
125
+ for i, val in enumerate(arr):
126
+ if val in self.vocab_:
127
+ out[i, self.vocab_[val]] = 1.0
128
+ elif self.handle_unknown == "error":
129
+ raise ValueError(f"Encountered unknown category: '{val}'")
130
+ return out
131
+
132
+ elif self.strategy == "ordinal":
133
+ out = np.zeros((num_samples, 1), dtype=np.int64)
134
+ for i, val in enumerate(arr):
135
+ if val in self.vocab_:
136
+ out[i, 0] = self.vocab_[val]
137
+ elif self.handle_unknown == "error":
138
+ raise ValueError(f"Encountered unknown category: '{val}'")
139
+ else:
140
+ out[i, 0] = self.unknown_value
141
+ return out
142
+
143
+ elif self.strategy == "hash":
144
+ out = np.zeros((num_samples, self.hash_dim), dtype=np.float32)
145
+ for i, val in enumerate(arr):
146
+ h = abs(hash(val)) % self.hash_dim
147
+ out[i, h] = 1.0
148
+ return out
149
+
150
+ else:
151
+ raise ValueError(f"Unknown categorical strategy '{self.strategy}'.")
152
+
153
+ def fit_transform(self, values: Sequence[Any]) -> np.ndarray:
154
+ return self.fit(values).transform(values)
155
+
156
+
157
+ class TabularEncoder:
158
+ r"""Column-wise encoder aggregating multiple numerical and categorical column encoders.
159
+
160
+ Args:
161
+ column_encoders: Optional dictionary mapping column names to :class:`NumericalEncoder`
162
+ or :class:`CategoricalEncoder` instances.
163
+ default_numerical_strategy: Strategy used for detected numerical columns without an explicit encoder.
164
+ (default: ``"standard"``)
165
+ default_categorical_strategy: Strategy used for detected categorical columns without an explicit encoder.
166
+ (default: ``"onehot"``)
167
+ """
168
+
169
+ def __init__(
170
+ self,
171
+ column_encoders: Optional[Dict[str, Union[NumericalEncoder, CategoricalEncoder]]] = None,
172
+ default_numerical_strategy: str = "standard",
173
+ default_categorical_strategy: str = "onehot",
174
+ ):
175
+ self.column_encoders = column_encoders or {}
176
+ self.default_numerical_strategy = default_numerical_strategy
177
+ self.default_categorical_strategy = default_categorical_strategy
178
+ self.fitted_encoders_: Dict[str, Union[NumericalEncoder, CategoricalEncoder]] = {}
179
+ self.column_order_: List[str] = []
180
+
181
+ def fit(self, df_or_dict: Any, columns: Optional[List[str]] = None):
182
+ columns = columns or _get_column_names(df_or_dict)
183
+ self.column_order_ = list(columns)
184
+ self.fitted_encoders_ = {}
185
+
186
+ for col in self.column_order_:
187
+ vals = _get_column_values(df_or_dict, col)
188
+ if col in self.column_encoders:
189
+ enc = self.column_encoders[col]
190
+ else:
191
+ if _is_numerical_series(vals):
192
+ enc = NumericalEncoder(strategy=self.default_numerical_strategy)
193
+ else:
194
+ enc = CategoricalEncoder(strategy=self.default_categorical_strategy)
195
+ enc.fit(vals)
196
+ self.fitted_encoders_[col] = enc
197
+
198
+ return self
199
+
200
+ def transform(self, df_or_dict: Any) -> np.ndarray:
201
+ parts = []
202
+ for col in self.column_order_:
203
+ vals = _get_column_values(df_or_dict, col)
204
+ enc = self.fitted_encoders_[col]
205
+ encoded = enc.transform(vals)
206
+ parts.append(encoded)
207
+
208
+ if not parts:
209
+ num_rows = len(_get_column_values(df_or_dict, list(self.column_encoders.keys())[0])) if self.column_encoders else 0
210
+ return np.zeros((num_rows, 0), dtype=np.float32)
211
+
212
+ return np.concatenate(parts, axis=1).astype(np.float32)
213
+
214
+ def fit_transform(self, df_or_dict: Any, columns: Optional[List[str]] = None) -> np.ndarray:
215
+ return self.fit(df_or_dict, columns=columns).transform(df_or_dict)
216
+
217
+
218
+ def _get_column_names(df_or_dict: Any) -> List[str]:
219
+ if hasattr(df_or_dict, "columns"):
220
+ return list(df_or_dict.columns)
221
+ elif isinstance(df_or_dict, dict):
222
+ return list(df_or_dict.keys())
223
+ raise TypeError(f"Expected pandas DataFrame or dictionary of columns, got {type(df_or_dict)}")
224
+
225
+
226
+ def _get_column_values(df_or_dict: Any, col: str) -> List[Any]:
227
+ if hasattr(df_or_dict, "__getitem__"):
228
+ series = df_or_dict[col]
229
+ if hasattr(series, "tolist"):
230
+ return series.tolist()
231
+ elif hasattr(series, "to_numpy"):
232
+ return series.to_numpy().tolist()
233
+ elif isinstance(series, np.ndarray):
234
+ return series.tolist()
235
+ return list(series)
236
+ raise TypeError(f"Cannot extract column '{col}' from object of type {type(df_or_dict)}")
237
+
238
+
239
+ def _is_numerical_series(vals: Sequence[Any]) -> bool:
240
+ count_num = 0
241
+ total = 0
242
+ for v in vals:
243
+ if v is None or (isinstance(v, float) and np.isnan(v)):
244
+ continue
245
+ total += 1
246
+ if isinstance(v, (int, float, np.integer, np.floating)) and not isinstance(v, bool):
247
+ count_num += 1
248
+ return total > 0 and (count_num / total) > 0.8