dewml 0.1.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 (335) hide show
  1. dew/__init__.py +143 -0
  2. dew/_model_types.py +9 -0
  3. dew/artifacts.py +117 -0
  4. dew/cache.py +65 -0
  5. dew/checkpoints/__init__.py +1740 -0
  6. dew/cli/__init__.py +1 -0
  7. dew/cli/config.py +180 -0
  8. dew/cli/gcloud.py +241 -0
  9. dew/cli/launch.py +599 -0
  10. dew/cli/main.py +136 -0
  11. dew/cli/ssh_config.py +69 -0
  12. dew/cli/tpu.py +912 -0
  13. dew/cli/tpu_setup.py +232 -0
  14. dew/config/__init__.py +1220 -0
  15. dew/config/sweep.py +114 -0
  16. dew/coordination.py +438 -0
  17. dew/data/__init__.py +57 -0
  18. dew/data/chat.py +733 -0
  19. dew/data/dataset.py +1772 -0
  20. dew/data/image_augmentation.py +154 -0
  21. dew/data/images.py +666 -0
  22. dew/data/online_loader.py +438 -0
  23. dew/data/preferences.py +165 -0
  24. dew/data/processors.py +58 -0
  25. dew/data/prompts.py +271 -0
  26. dew/data/providers.py +605 -0
  27. dew/data/rows.py +99 -0
  28. dew/data/sources/av_utils.py +118 -0
  29. dew/data/sources/hf.py +241 -0
  30. dew/data/sources/hf_stream.py +298 -0
  31. dew/data/sources/pytorch.py +60 -0
  32. dew/data/sources/text.py +574 -0
  33. dew/data/sources/tfds.py +281 -0
  34. dew/data/streaming.py +93 -0
  35. dew/data/text.py +199 -0
  36. dew/data/tokens.py +406 -0
  37. dew/data/video.py +172 -0
  38. dew/decision/__init__.py +115 -0
  39. dew/decision/calibration.py +248 -0
  40. dew/decision/clef.py +162 -0
  41. dew/decision/config.py +191 -0
  42. dew/decision/data.py +251 -0
  43. dew/decision/head.py +333 -0
  44. dew/decision/images.py +53 -0
  45. dew/decision/laya.py +239 -0
  46. dew/decision/layout.py +505 -0
  47. dew/decision/metrics.py +129 -0
  48. dew/decision/model.py +107 -0
  49. dew/decision/objective.py +488 -0
  50. dew/decision/questions.py +343 -0
  51. dew/decision/scoring.py +155 -0
  52. dew/decision/task.py +627 -0
  53. dew/diffusion/__init__.py +72 -0
  54. dew/diffusion/block.py +392 -0
  55. dew/diffusion/discrete.py +377 -0
  56. dew/diffusion/presets.py +331 -0
  57. dew/diffusion/process.py +224 -0
  58. dew/diffusion/schedules/__init__.py +14 -0
  59. dew/diffusion/schedules/common.py +132 -0
  60. dew/diffusion/schedules/cosine.py +42 -0
  61. dew/diffusion/schedules/discrete.py +58 -0
  62. dew/diffusion/schedules/flow.py +87 -0
  63. dew/diffusion/schedules/karras.py +57 -0
  64. dew/diffusion/schedules/linear.py +30 -0
  65. dew/diffusion/schedules/source.py +100 -0
  66. dew/diffusion/schedules/source_grids.py +623 -0
  67. dew/diffusion/schedules/source_policy.py +626 -0
  68. dew/diffusion/schedules/sqrt.py +36 -0
  69. dew/diffusion/transforms.py +347 -0
  70. dew/eval/__init__.py +31 -0
  71. dew/eval/__main__.py +26 -0
  72. dew/eval/common.py +113 -0
  73. dew/eval/fid.py +282 -0
  74. dew/eval/harness.py +288 -0
  75. dew/eval/images.py +147 -0
  76. dew/eval/inception.py +207 -0
  77. dew/eval/lpips.py +144 -0
  78. dew/eval/psnr.py +60 -0
  79. dew/eval/ssim.py +128 -0
  80. dew/files.py +89 -0
  81. dew/inference/__init__.py +48 -0
  82. dew/inference/banks.py +769 -0
  83. dew/inference/clients.py +628 -0
  84. dew/inference/nccl.py +320 -0
  85. dew/inference/pages.py +137 -0
  86. dew/inference/pipeline.py +209 -0
  87. dew/inference/projections.py +86 -0
  88. dew/inference/rollouts.py +595 -0
  89. dew/inference/serving.py +896 -0
  90. dew/inference/serving_kernel.py +513 -0
  91. dew/inference/tasks.py +748 -0
  92. dew/inputs/__init__.py +162 -0
  93. dew/inputs/diffusion.py +690 -0
  94. dew/inputs/encoders.py +471 -0
  95. dew/interop/__init__.py +59 -0
  96. dew/interop/codecs.py +1579 -0
  97. dew/interop/components.py +56 -0
  98. dew/interop/config_records.py +45 -0
  99. dew/interop/dduf.py +72 -0
  100. dew/interop/decoder_families.py +123 -0
  101. dew/interop/decoder_parts.py +1342 -0
  102. dew/interop/diffusion.py +951 -0
  103. dew/interop/diffusion_gemma.py +300 -0
  104. dew/interop/families/__init__.py +4 -0
  105. dew/interop/families/bloom.py +88 -0
  106. dew/interop/families/deepseek.py +1192 -0
  107. dew/interop/families/deepseek_v41.py +393 -0
  108. dew/interop/families/falcon.py +105 -0
  109. dew/interop/families/gemma.py +650 -0
  110. dew/interop/families/glm.py +553 -0
  111. dew/interop/families/gpt2.py +331 -0
  112. dew/interop/families/gpt_bigcode.py +88 -0
  113. dew/interop/families/gpt_neox.py +179 -0
  114. dew/interop/families/gpt_oss.py +102 -0
  115. dew/interop/families/kimi.py +473 -0
  116. dew/interop/families/llama.py +155 -0
  117. dew/interop/families/llama4.py +170 -0
  118. dew/interop/families/masked_diffusion.py +371 -0
  119. dew/interop/families/modernbert.py +222 -0
  120. dew/interop/families/nemotron_h.py +204 -0
  121. dew/interop/families/olmo.py +69 -0
  122. dew/interop/families/opt.py +98 -0
  123. dew/interop/families/phi.py +72 -0
  124. dew/interop/families/phi3.py +128 -0
  125. dew/interop/families/qwen.py +481 -0
  126. dew/interop/families/starcoder2.py +87 -0
  127. dew/interop/flaxdiff.py +290 -0
  128. dew/interop/generation_config.py +565 -0
  129. dew/interop/gguf.py +224 -0
  130. dew/interop/harbor.py +704 -0
  131. dew/interop/hf_decoders.py +927 -0
  132. dew/interop/hub.py +24 -0
  133. dew/interop/inception_fid.py +248 -0
  134. dew/interop/mamba2.py +242 -0
  135. dew/interop/pickles.py +117 -0
  136. dew/interop/pipeline_assembly.py +985 -0
  137. dew/interop/pretrained.py +1312 -0
  138. dew/interop/processors.py +598 -0
  139. dew/interop/safetensors_io.py +434 -0
  140. dew/interop/single_file.py +311 -0
  141. dew/interop/sources.py +195 -0
  142. dew/interop/streaming.py +249 -0
  143. dew/interop/torchax_fallback.py +366 -0
  144. dew/interop/verify.py +420 -0
  145. dew/interop/weights.py +189 -0
  146. dew/io.py +66 -0
  147. dew/logging.py +101 -0
  148. dew/lora.py +907 -0
  149. dew/nn/__init__.py +0 -0
  150. dew/nn/activations.py +89 -0
  151. dew/nn/attention.py +2225 -0
  152. dew/nn/attention_residuals.py +90 -0
  153. dew/nn/attention_sinks.py +53 -0
  154. dew/nn/audio.py +770 -0
  155. dew/nn/autoencoders/__init__.py +5 -0
  156. dew/nn/autoencoders/api.py +154 -0
  157. dew/nn/autoencoders/dc_ae.py +513 -0
  158. dew/nn/autoencoders/flux2.py +100 -0
  159. dew/nn/autoencoders/kl.py +108 -0
  160. dew/nn/autoencoders/pretrained.py +94 -0
  161. dew/nn/autoencoders/qwen_image.py +389 -0
  162. dew/nn/autoencoders/rae.py +516 -0
  163. dew/nn/autoencoders/sd_vae.py +64 -0
  164. dew/nn/autoencoders/vae.py +434 -0
  165. dew/nn/autoencoders/wan.py +509 -0
  166. dew/nn/backbones/__init__.py +30 -0
  167. dew/nn/backbones/causal_transformer.py +1942 -0
  168. dew/nn/backbones/decoder_block.py +989 -0
  169. dew/nn/backbones/decoder_stack.py +537 -0
  170. dew/nn/backbones/dit.py +119 -0
  171. dew/nn/backbones/edm2.py +162 -0
  172. dew/nn/backbones/flux.py +191 -0
  173. dew/nn/backbones/flux2.py +199 -0
  174. dew/nn/backbones/jepa.py +220 -0
  175. dew/nn/backbones/joint.py +296 -0
  176. dew/nn/backbones/layer_plan.py +192 -0
  177. dew/nn/backbones/mmdit.py +303 -0
  178. dew/nn/backbones/qwen_image.py +261 -0
  179. dew/nn/backbones/sd3.py +165 -0
  180. dew/nn/backbones/ssm_dit.py +83 -0
  181. dew/nn/backbones/unet.py +135 -0
  182. dew/nn/backbones/unet3d.py +105 -0
  183. dew/nn/backbones/unet_condition.py +312 -0
  184. dew/nn/backbones/uvit.py +261 -0
  185. dew/nn/backbones/video_dit.py +63 -0
  186. dew/nn/backbones/wan.py +250 -0
  187. dew/nn/backbones/z_image.py +235 -0
  188. dew/nn/blocks.py +285 -0
  189. dew/nn/conv.py +251 -0
  190. dew/nn/deepseek_v4.py +957 -0
  191. dew/nn/diffusion_gemma.py +357 -0
  192. dew/nn/dit.py +666 -0
  193. dew/nn/dsa_kpool.py +450 -0
  194. dew/nn/dspark.py +190 -0
  195. dew/nn/engram.py +326 -0
  196. dew/nn/fake_quant.py +117 -0
  197. dew/nn/gemma3n.py +225 -0
  198. dew/nn/gemma4_moe.py +99 -0
  199. dew/nn/gpt_oss.py +122 -0
  200. dew/nn/hyper_connections.py +206 -0
  201. dew/nn/inputs.py +720 -0
  202. dew/nn/kda.py +283 -0
  203. dew/nn/kernels/__init__.py +16 -0
  204. dew/nn/kernels/decode_attention.py +159 -0
  205. dew/nn/kernels/delta_rule.py +114 -0
  206. dew/nn/kernels/generation.py +96 -0
  207. dew/nn/kernels/grouped_matmul.py +181 -0
  208. dew/nn/kernels/ragged_dot.py +444 -0
  209. dew/nn/kernels/ssd.py +350 -0
  210. dew/nn/kv_cache.py +569 -0
  211. dew/nn/linear.py +671 -0
  212. dew/nn/llama4.py +104 -0
  213. dew/nn/mixer_base.py +125 -0
  214. dew/nn/mixers/__init__.py +4 -0
  215. dew/nn/mixers/attention.py +980 -0
  216. dew/nn/mixers/gated_delta_net.py +50 -0
  217. dew/nn/mixers/mamba2.py +592 -0
  218. dew/nn/mixers/mlp.py +88 -0
  219. dew/nn/mla.py +783 -0
  220. dew/nn/mobilenet.py +431 -0
  221. dew/nn/moe.py +1150 -0
  222. dew/nn/mp.py +171 -0
  223. dew/nn/multimodal.py +494 -0
  224. dew/nn/precision.py +252 -0
  225. dew/nn/protocols.py +304 -0
  226. dew/nn/rope.py +384 -0
  227. dew/nn/safety.py +56 -0
  228. dew/nn/scan_orders.py +148 -0
  229. dew/nn/scatter.py +13 -0
  230. dew/nn/sharding.py +639 -0
  231. dew/nn/sparse_selection.py +140 -0
  232. dew/nn/ssm.py +276 -0
  233. dew/nn/text_encoders.py +858 -0
  234. dew/nn/vision/__init__.py +137 -0
  235. dew/nn/vision/common.py +118 -0
  236. dew/nn/vision/deepseek_v41.py +212 -0
  237. dew/nn/vision/gemma3n.py +235 -0
  238. dew/nn/vision/gemma4.py +394 -0
  239. dew/nn/vision/llama4.py +296 -0
  240. dew/nn/vision/qwen35.py +342 -0
  241. dew/nn/vision/siglip.py +258 -0
  242. dew/objectives/__init__.py +4 -0
  243. dew/objectives/base.py +1055 -0
  244. dew/objectives/diffusion/__init__.py +28 -0
  245. dew/objectives/diffusion/adversarial.py +330 -0
  246. dew/objectives/diffusion/alignment.py +192 -0
  247. dew/objectives/diffusion/block.py +439 -0
  248. dew/objectives/diffusion/config.py +511 -0
  249. dew/objectives/diffusion/consistency.py +446 -0
  250. dew/objectives/diffusion/end_to_end.py +272 -0
  251. dew/objectives/diffusion/few_step.py +317 -0
  252. dew/objectives/diffusion/guidance_distillation.py +134 -0
  253. dew/objectives/diffusion/masked.py +242 -0
  254. dew/objectives/diffusion/objective.py +867 -0
  255. dew/objectives/distillation.py +256 -0
  256. dew/objectives/jepa/__init__.py +16 -0
  257. dew/objectives/jepa/config.py +86 -0
  258. dew/objectives/jepa/masking.py +148 -0
  259. dew/objectives/jepa/objective.py +258 -0
  260. dew/objectives/jepa/probes.py +149 -0
  261. dew/objectives/lm/__init__.py +4 -0
  262. dew/objectives/lm/chunked.py +862 -0
  263. dew/objectives/lm/config.py +288 -0
  264. dew/objectives/lm/objective.py +1276 -0
  265. dew/objectives/rl/__init__.py +93 -0
  266. dew/objectives/rl/episodes.py +737 -0
  267. dew/objectives/rl/flow.py +442 -0
  268. dew/objectives/rl/grpo.py +363 -0
  269. dew/objectives/rl/journal.py +158 -0
  270. dew/objectives/rl/ppo.py +325 -0
  271. dew/objectives/rl/preference.py +117 -0
  272. dew/objectives/rl/records.py +70 -0
  273. dew/objectives/rl/rollout.py +154 -0
  274. dew/objectives/rl/scheduler.py +572 -0
  275. dew/objectives/rl/sessions.py +674 -0
  276. dew/objectives/rl/sources.py +372 -0
  277. dew/objectives/rl/verl.py +314 -0
  278. dew/objectives/supervised.py +138 -0
  279. dew/pool.py +123 -0
  280. dew/position.py +98 -0
  281. dew/py.typed +0 -0
  282. dew/records.py +138 -0
  283. dew/registry.py +1250 -0
  284. dew/rl/__init__.py +46 -0
  285. dew/rl/_sandbox_exec.py +28 -0
  286. dew/rl/advantage.py +170 -0
  287. dew/rl/sandbox.py +670 -0
  288. dew/rl/surrogate.py +363 -0
  289. dew/sampling/__init__.py +72 -0
  290. dew/sampling/decoding.py +968 -0
  291. dew/sampling/flow.py +216 -0
  292. dew/sampling/guidance.py +379 -0
  293. dew/sampling/guided.py +164 -0
  294. dew/sampling/pipelines.py +854 -0
  295. dew/sampling/sample.py +144 -0
  296. dew/sampling/solvers/__init__.py +25 -0
  297. dew/sampling/solvers/brownian.py +176 -0
  298. dew/sampling/solvers/common.py +144 -0
  299. dew/sampling/solvers/dpm.py +388 -0
  300. dew/sampling/solvers/gaussian.py +230 -0
  301. dew/sampling/solvers/sigma.py +300 -0
  302. dew/sampling/solvers/unipc.py +213 -0
  303. dew/sampling/strategies.py +899 -0
  304. dew/sampling/text.py +953 -0
  305. dew/sampling/vocabulary.py +188 -0
  306. dew/telemetry/__init__.py +0 -0
  307. dew/telemetry/devices.py +160 -0
  308. dew/telemetry/instrumentation.py +453 -0
  309. dew/telemetry/peaks.py +51 -0
  310. dew/telemetry/profile.py +341 -0
  311. dew/telemetry/records.py +149 -0
  312. dew/training/__init__.py +25 -0
  313. dew/training/display.py +540 -0
  314. dew/training/distributed.py +905 -0
  315. dew/training/evaluation.py +493 -0
  316. dew/training/execution.py +429 -0
  317. dew/training/host.py +255 -0
  318. dew/training/memory.py +341 -0
  319. dew/training/narrow.py +99 -0
  320. dew/training/optim.py +896 -0
  321. dew/training/posthoc.py +93 -0
  322. dew/training/quantization.py +870 -0
  323. dew/training/rungs.py +70 -0
  324. dew/training/runtime.py +323 -0
  325. dew/training/selection.py +50 -0
  326. dew/training/state.py +82 -0
  327. dew/training/tracker.py +661 -0
  328. dew/training/trainer.py +2266 -0
  329. dew/training/transaction.py +486 -0
  330. dewml-0.1.0.dist-info/METADATA +1594 -0
  331. dewml-0.1.0.dist-info/RECORD +335 -0
  332. dewml-0.1.0.dist-info/WHEEL +5 -0
  333. dewml-0.1.0.dist-info/entry_points.txt +2 -0
  334. dewml-0.1.0.dist-info/licenses/LICENSE +21 -0
  335. dewml-0.1.0.dist-info/top_level.txt +1 -0
dew/__init__.py ADDED
@@ -0,0 +1,143 @@
1
+ """Dew: one objective, one trainer.
2
+
3
+ Each name exported here is imported from its own module when you first
4
+ access it, not at `import dew`. So `import dew.training` loads only the
5
+ training layer, with no modality, encoder or tracker backend. Importing
6
+ `dew` opens no JAX backend and loads no optional dependency; encoders,
7
+ decoders and datasets load what they need when they are built.
8
+
9
+ `import dew` does set two XLA flags before the backend opens.
10
+ `--xla_allow_excess_precision=false` makes XLA round values declared in a
11
+ narrow dtype such as bf16 where the program rounds them. When JAX's CUDA
12
+ plugin is installed, `--xla_gpu_enable_allocator_spatial_partitioning=false`
13
+ keeps a preallocated GPU pool in one piece for a step's temporaries. If you
14
+ set either flag yourself in XLA_FLAGS, your value is kept. If the JAX
15
+ backend has already opened, the flags cannot take effect, and Dew logs a
16
+ warning.
17
+ """
18
+
19
+ from collections.abc import Callable
20
+ from importlib import import_module
21
+ from typing import TYPE_CHECKING
22
+
23
+ from dew.logging import configure as _configure_logging
24
+ from dew.telemetry.devices import (
25
+ keep_roundings as _keep_roundings,
26
+ unpartition_gpu_pool as _unpartition_gpu_pool,
27
+ )
28
+
29
+ _configure_logging()
30
+ _keep_roundings()
31
+ _unpartition_gpu_pool()
32
+
33
+ if TYPE_CHECKING: # the surface above, with its types, for checkers and editors
34
+ from dew.artifacts import ImageGrid, Representations, TextSamples, TokenScores, VideoGrid
35
+ from dew.data import Dataset
36
+ from dew.diffusion import Process
37
+ from dew.eval import Mean
38
+ from dew.inference import pipeline
39
+ from dew.inputs import Condition, Field, InputSpec
40
+ from dew.objectives import Objective
41
+ from dew.objectives.base import Aux, EMASpec, Step
42
+ from dew.objectives.supervised import Supervised
43
+ from dew.sampling import CFG, sample
44
+ from dew.telemetry.profile import Profiler
45
+ from dew.training import (
46
+ Best,
47
+ Checkpoints,
48
+ EvalSuite,
49
+ Evaluation,
50
+ Keep,
51
+ Layout,
52
+ LocalTracker,
53
+ MeshSpec,
54
+ MLflowTracker,
55
+ Plateau,
56
+ ProfileWindow,
57
+ TensorBoardTracker,
58
+ Tracker,
59
+ Trackers,
60
+ Trainer,
61
+ TrainState,
62
+ WandbTracker,
63
+ )
64
+
65
+ __version__ = "0.1.0"
66
+
67
+ _EXPORTS = {
68
+ "Best": "dew.training", "Keep": "dew.training", "Plateau": "dew.training", "EvalSuite": "dew.training",
69
+ "Trainer": "dew.training", "TrainState": "dew.training", "Step": "dew.training",
70
+ "Aux": "dew.training", "EMASpec": "dew.training", "MeshSpec": "dew.training",
71
+ "Layout": "dew.training", "Checkpoints": "dew.training", "Tracker": "dew.training",
72
+ "WandbTracker": "dew.training", "LocalTracker": "dew.training", "Trackers": "dew.training",
73
+ "MLflowTracker": "dew.training", "TensorBoardTracker": "dew.training",
74
+ "Evaluation": "dew.training",
75
+ "ProfileWindow": "dew.training",
76
+ "Objective": "dew.objectives", "Supervised": "dew.objectives.supervised",
77
+ "Dataset": "dew.data",
78
+ "Process": "dew.diffusion",
79
+ "Mean": "dew.eval",
80
+ "InputSpec": "dew.inputs", "Field": "dew.inputs", "Condition": "dew.inputs",
81
+ "sample": "dew.sampling", "CFG": "dew.sampling",
82
+ "pipeline": "dew.inference",
83
+ "Profiler": "dew.telemetry.profile",
84
+ "ImageGrid": "dew.artifacts", "VideoGrid": "dew.artifacts",
85
+ "TextSamples": "dew.artifacts", "Representations": "dew.artifacts",
86
+ "TokenScores": "dew.artifacts",
87
+ }
88
+
89
+
90
+ def __getattr__(name: str) -> type | Callable:
91
+ module = _EXPORTS.get(name)
92
+ if module is None:
93
+ raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
94
+ return getattr(import_module(module), name)
95
+
96
+
97
+ def __dir__() -> list[str]:
98
+ return list(__all__)
99
+
100
+
101
+ # Written out, not derived from _EXPORTS, so a type checker, an editor and
102
+ # `from dew import *` can all read the public surface without running the
103
+ # lazy lookup above.
104
+ __all__ = [
105
+ "CFG",
106
+ "Aux",
107
+ "Best",
108
+ "Checkpoints",
109
+ "Condition",
110
+ "Dataset",
111
+ "EMASpec",
112
+ "EvalSuite",
113
+ "Evaluation",
114
+ "Field",
115
+ "ImageGrid",
116
+ "InputSpec",
117
+ "Keep",
118
+ "Layout",
119
+ "LocalTracker",
120
+ "MLflowTracker",
121
+ "Mean",
122
+ "MeshSpec",
123
+ "Objective",
124
+ "Plateau",
125
+ "Process",
126
+ "ProfileWindow",
127
+ "Profiler",
128
+ "Representations",
129
+ "Step",
130
+ "Supervised",
131
+ "TensorBoardTracker",
132
+ "TextSamples",
133
+ "TokenScores",
134
+ "Tracker",
135
+ "Trackers",
136
+ "TrainState",
137
+ "Trainer",
138
+ "VideoGrid",
139
+ "WandbTracker",
140
+ "__version__",
141
+ "pipeline",
142
+ "sample",
143
+ ]
dew/_model_types.py ADDED
@@ -0,0 +1,9 @@
1
+ """Model-type spellings shared by the Qwen decoder, wrapper and vision maps.
2
+
3
+ Vision translation reads these without importing the interop package, whose
4
+ public entry point imports the vision classes itself.
5
+ """
6
+
7
+ QWEN35_TYPES = ('qwen3_5', 'qwen3_5_moe')
8
+ QWEN35_TEXT_TYPES = tuple(f'{name}_text' for name in QWEN35_TYPES)
9
+ _QWEN35_VISION_TYPES = (*QWEN35_TYPES, *(f'{name}_vision' for name in QWEN35_TYPES))
dew/artifacts.py ADDED
@@ -0,0 +1,117 @@
1
+ """Typed values that an objective's evaluation produces.
2
+
3
+ An objective returns these from its scoring and preview hooks. A metric
4
+ reads the scoring artifact of the type it expects, and a tracker renders
5
+ each preview according to its type. The array fields are pytree leaves, so
6
+ they can pass through jit, while the optional captions and decoded text stay
7
+ on the host as static metadata.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ from typing import TYPE_CHECKING
13
+
14
+ import jax
15
+ import numpy as np
16
+ from flax import struct
17
+ from jax.typing import ArrayLike
18
+ from numpy.typing import NDArray
19
+
20
+ if TYPE_CHECKING:
21
+ from _typeshed import DataclassInstance
22
+
23
+
24
+ @struct.dataclass
25
+ class ImageGrid:
26
+ """Images in [-1, 1], `[N, H, W, C]`, with the text each was conditioned on
27
+ where there was any."""
28
+ images: jax.Array
29
+ captions: tuple[str, ...] = struct.field(pytree_node=False, default=())
30
+
31
+
32
+ @struct.dataclass
33
+ class VideoGrid:
34
+ """Clips in [-1, 1], `[N, T, H, W, C]`, with the text each was conditioned on
35
+ where there was any."""
36
+ videos: jax.Array
37
+ captions: tuple[str, ...] = struct.field(pytree_node=False, default=())
38
+
39
+
40
+ @struct.dataclass
41
+ class TextSamples:
42
+ """Generated token rows, with optional decoded preview text and prompt."""
43
+ tokens: jax.Array | np.ndarray
44
+ prompt: str = struct.field(pytree_node=False, default="")
45
+ texts: tuple[str, ...] = struct.field(pytree_node=False, default=())
46
+
47
+
48
+ @struct.dataclass
49
+ class Representations:
50
+ """Encoder outputs `[N, D]` and the labels of the records they came from,
51
+ for a probe to score."""
52
+ features: jax.Array
53
+ labels: jax.Array
54
+
55
+
56
+ @struct.dataclass
57
+ class TokenScores:
58
+ """Teacher-forced per-token losses `[N, L]` and the weight of each target.
59
+
60
+ A weight is 1 where the target counts and 0 where it is padding or a
61
+ document's first token. A perplexity is exp of the weighted mean loss
62
+ over a whole pass, so a batch with no counted target adds nothing to it.
63
+ """
64
+ losses: jax.Array
65
+ weights: jax.Array
66
+ correct: jax.Array
67
+ """Per-token top-1 correctness, from the same logits that produced the losses."""
68
+
69
+
70
+ @struct.dataclass
71
+ class Decisions:
72
+ """A decision model's probabilities `[N, Q, K]` over the options of each
73
+ row's questions, the real options `[N, Q, K]`, each question's right option
74
+ `[N, Q]`, which questions `[N, Q]` ask a score, whose options are ordered
75
+ levels, and which `[N, Q]` have a known answer to score."""
76
+ probabilities: jax.Array
77
+ options: jax.Array
78
+ labels: jax.Array
79
+ ordinal: jax.Array
80
+ scored: jax.Array
81
+
82
+
83
+ type Artifact = DataclassInstance
84
+ """What an objective's scoring and preview hooks return: a dataclass whose
85
+ per-row fields lead with the batch's rows, which a validation pass cuts to
86
+ the real ones. Dew's own are the classes above. A package's objective may
87
+ score into a dataclass of its own, which its metrics read by type
88
+ (`Metric.reads`); a metric picks exactly one scoring artifact, and previews
89
+ never satisfy metrics. A tracker shows a preview of Dew's types, so a
90
+ package shows its own as one of them, a spike raster as an `ImageGrid`."""
91
+
92
+ type Artifacts = Artifact | tuple[Artifact, ...]
93
+ """One artifact, or several."""
94
+
95
+
96
+ def uint8_pixels(images: ArrayLike) -> NDArray[np.uint8]:
97
+ """Convert [-1, 1] pixels, as `ImageGrid` and `VideoGrid` hold them, to uint8 in [0, 255].
98
+
99
+ Each value maps to its nearest level, computed as `(x + 1) * 127.5` in
100
+ float32 and rounded half to even (`np.rint`). The result is clipped,
101
+ because a sample can leave the range. Metrics score these bytes and
102
+ trackers preview them, so both see the same image.
103
+ """
104
+ levels = np.rint((np.asarray(images, np.float32) + 1.0) * 127.5)
105
+ return np.clip(levels, 0, 255).astype(np.uint8)
106
+
107
+
108
+ __all__ = [
109
+ "Artifact",
110
+ "Decisions",
111
+ "ImageGrid",
112
+ "Representations",
113
+ "TextSamples",
114
+ "TokenScores",
115
+ "VideoGrid",
116
+ "uint8_pixels",
117
+ ]
dew/cache.py ADDED
@@ -0,0 +1,65 @@
1
+ """Where Dew keeps what it caches on disk, and JAX's persistent compilation cache.
2
+
3
+ Kept apart from the telemetry and the loaders that read it, so the
4
+ interop, config and inference modules take their cache paths without
5
+ importing the FLOP accounting.
6
+ """
7
+
8
+ import os
9
+ import sys
10
+
11
+ import jax
12
+
13
+
14
+ def dew_cache_dir() -> str:
15
+ """Dew's cache directory: `$XDG_CACHE_HOME/dew`, else ~/.cache/dew."""
16
+ return os.path.expanduser(
17
+ os.path.join(os.environ.get("XDG_CACHE_HOME") or os.path.join("~", ".cache"), "dew")
18
+ )
19
+
20
+
21
+ def default_compilation_cache_dir() -> str:
22
+ """Where compiled executables go unless a run names somewhere else.
23
+
24
+ The directory JAX is configured with (`jax_compilation_cache_dir`, which
25
+ JAX_COMPILATION_CACHE_DIR sets) when there is one, so a machine keeps one
26
+ cache for every entry point. Otherwise Python minors have separate
27
+ defaults: jax 0.11.2 compresses with Python 3.14's stdlib zstd but names
28
+ the codec "zlib" in the key, which says "zstandard" only for the
29
+ zstandard package (`jax._src.compilation_cache.get_cache_key`), so an
30
+ older interpreter sharing the directory would read those bytes with the
31
+ wrong codec. Explicit paths passed to enable_compilation_cache remain
32
+ unchanged.
33
+ """
34
+ if jax.config.jax_compilation_cache_dir:
35
+ return jax.config.jax_compilation_cache_dir
36
+ return os.path.join(dew_cache_dir(), 'xla', f"python{sys.version_info.major}.{sys.version_info.minor}")
37
+
38
+
39
+ def enable_compilation_cache(path: str):
40
+ """Persist compiled executables so restarts skip XLA compilation.
41
+
42
+ The dominant cost of a restart-heavy TPU workflow, where every run otherwise
43
+ recompiles the same step function from scratch.
44
+ """
45
+ os.makedirs(path, exist_ok=True)
46
+ jax.config.update('jax_compilation_cache_dir', path)
47
+ # Defaults skip small/fast compilations; a training step is neither, and
48
+ # caching everything keeps startup predictable.
49
+ jax.config.update('jax_persistent_cache_min_entry_size_bytes', -1)
50
+ jax.config.update('jax_persistent_cache_min_compile_time_secs', 0.0)
51
+
52
+
53
+ def persist_compilations() -> None:
54
+ """Point XLA at the on-disk executable cache, unless a directory is set.
55
+
56
+ A loaded task compiles for seconds the first time its shapes are seen
57
+ (a minute for a text-to-image sample on an A100), and a serving process
58
+ restarts. Training turns the same cache on in `prepare_process`; a task
59
+ turns it on where it is loaded, `dew.pipeline` or a saved run's record
60
+ (`dew.inference.tasks.run_record`). Reading the setting is what makes it
61
+ idempotent and what leaves a trainer's own directory, or a caller's, alone.
62
+ """
63
+ if jax.config.jax_compilation_cache_dir:
64
+ return
65
+ enable_compilation_cache(default_compilation_cache_dir())