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.
- dew/__init__.py +143 -0
- dew/_model_types.py +9 -0
- dew/artifacts.py +117 -0
- dew/cache.py +65 -0
- dew/checkpoints/__init__.py +1740 -0
- dew/cli/__init__.py +1 -0
- dew/cli/config.py +180 -0
- dew/cli/gcloud.py +241 -0
- dew/cli/launch.py +599 -0
- dew/cli/main.py +136 -0
- dew/cli/ssh_config.py +69 -0
- dew/cli/tpu.py +912 -0
- dew/cli/tpu_setup.py +232 -0
- dew/config/__init__.py +1220 -0
- dew/config/sweep.py +114 -0
- dew/coordination.py +438 -0
- dew/data/__init__.py +57 -0
- dew/data/chat.py +733 -0
- dew/data/dataset.py +1772 -0
- dew/data/image_augmentation.py +154 -0
- dew/data/images.py +666 -0
- dew/data/online_loader.py +438 -0
- dew/data/preferences.py +165 -0
- dew/data/processors.py +58 -0
- dew/data/prompts.py +271 -0
- dew/data/providers.py +605 -0
- dew/data/rows.py +99 -0
- dew/data/sources/av_utils.py +118 -0
- dew/data/sources/hf.py +241 -0
- dew/data/sources/hf_stream.py +298 -0
- dew/data/sources/pytorch.py +60 -0
- dew/data/sources/text.py +574 -0
- dew/data/sources/tfds.py +281 -0
- dew/data/streaming.py +93 -0
- dew/data/text.py +199 -0
- dew/data/tokens.py +406 -0
- dew/data/video.py +172 -0
- dew/decision/__init__.py +115 -0
- dew/decision/calibration.py +248 -0
- dew/decision/clef.py +162 -0
- dew/decision/config.py +191 -0
- dew/decision/data.py +251 -0
- dew/decision/head.py +333 -0
- dew/decision/images.py +53 -0
- dew/decision/laya.py +239 -0
- dew/decision/layout.py +505 -0
- dew/decision/metrics.py +129 -0
- dew/decision/model.py +107 -0
- dew/decision/objective.py +488 -0
- dew/decision/questions.py +343 -0
- dew/decision/scoring.py +155 -0
- dew/decision/task.py +627 -0
- dew/diffusion/__init__.py +72 -0
- dew/diffusion/block.py +392 -0
- dew/diffusion/discrete.py +377 -0
- dew/diffusion/presets.py +331 -0
- dew/diffusion/process.py +224 -0
- dew/diffusion/schedules/__init__.py +14 -0
- dew/diffusion/schedules/common.py +132 -0
- dew/diffusion/schedules/cosine.py +42 -0
- dew/diffusion/schedules/discrete.py +58 -0
- dew/diffusion/schedules/flow.py +87 -0
- dew/diffusion/schedules/karras.py +57 -0
- dew/diffusion/schedules/linear.py +30 -0
- dew/diffusion/schedules/source.py +100 -0
- dew/diffusion/schedules/source_grids.py +623 -0
- dew/diffusion/schedules/source_policy.py +626 -0
- dew/diffusion/schedules/sqrt.py +36 -0
- dew/diffusion/transforms.py +347 -0
- dew/eval/__init__.py +31 -0
- dew/eval/__main__.py +26 -0
- dew/eval/common.py +113 -0
- dew/eval/fid.py +282 -0
- dew/eval/harness.py +288 -0
- dew/eval/images.py +147 -0
- dew/eval/inception.py +207 -0
- dew/eval/lpips.py +144 -0
- dew/eval/psnr.py +60 -0
- dew/eval/ssim.py +128 -0
- dew/files.py +89 -0
- dew/inference/__init__.py +48 -0
- dew/inference/banks.py +769 -0
- dew/inference/clients.py +628 -0
- dew/inference/nccl.py +320 -0
- dew/inference/pages.py +137 -0
- dew/inference/pipeline.py +209 -0
- dew/inference/projections.py +86 -0
- dew/inference/rollouts.py +595 -0
- dew/inference/serving.py +896 -0
- dew/inference/serving_kernel.py +513 -0
- dew/inference/tasks.py +748 -0
- dew/inputs/__init__.py +162 -0
- dew/inputs/diffusion.py +690 -0
- dew/inputs/encoders.py +471 -0
- dew/interop/__init__.py +59 -0
- dew/interop/codecs.py +1579 -0
- dew/interop/components.py +56 -0
- dew/interop/config_records.py +45 -0
- dew/interop/dduf.py +72 -0
- dew/interop/decoder_families.py +123 -0
- dew/interop/decoder_parts.py +1342 -0
- dew/interop/diffusion.py +951 -0
- dew/interop/diffusion_gemma.py +300 -0
- dew/interop/families/__init__.py +4 -0
- dew/interop/families/bloom.py +88 -0
- dew/interop/families/deepseek.py +1192 -0
- dew/interop/families/deepseek_v41.py +393 -0
- dew/interop/families/falcon.py +105 -0
- dew/interop/families/gemma.py +650 -0
- dew/interop/families/glm.py +553 -0
- dew/interop/families/gpt2.py +331 -0
- dew/interop/families/gpt_bigcode.py +88 -0
- dew/interop/families/gpt_neox.py +179 -0
- dew/interop/families/gpt_oss.py +102 -0
- dew/interop/families/kimi.py +473 -0
- dew/interop/families/llama.py +155 -0
- dew/interop/families/llama4.py +170 -0
- dew/interop/families/masked_diffusion.py +371 -0
- dew/interop/families/modernbert.py +222 -0
- dew/interop/families/nemotron_h.py +204 -0
- dew/interop/families/olmo.py +69 -0
- dew/interop/families/opt.py +98 -0
- dew/interop/families/phi.py +72 -0
- dew/interop/families/phi3.py +128 -0
- dew/interop/families/qwen.py +481 -0
- dew/interop/families/starcoder2.py +87 -0
- dew/interop/flaxdiff.py +290 -0
- dew/interop/generation_config.py +565 -0
- dew/interop/gguf.py +224 -0
- dew/interop/harbor.py +704 -0
- dew/interop/hf_decoders.py +927 -0
- dew/interop/hub.py +24 -0
- dew/interop/inception_fid.py +248 -0
- dew/interop/mamba2.py +242 -0
- dew/interop/pickles.py +117 -0
- dew/interop/pipeline_assembly.py +985 -0
- dew/interop/pretrained.py +1312 -0
- dew/interop/processors.py +598 -0
- dew/interop/safetensors_io.py +434 -0
- dew/interop/single_file.py +311 -0
- dew/interop/sources.py +195 -0
- dew/interop/streaming.py +249 -0
- dew/interop/torchax_fallback.py +366 -0
- dew/interop/verify.py +420 -0
- dew/interop/weights.py +189 -0
- dew/io.py +66 -0
- dew/logging.py +101 -0
- dew/lora.py +907 -0
- dew/nn/__init__.py +0 -0
- dew/nn/activations.py +89 -0
- dew/nn/attention.py +2225 -0
- dew/nn/attention_residuals.py +90 -0
- dew/nn/attention_sinks.py +53 -0
- dew/nn/audio.py +770 -0
- dew/nn/autoencoders/__init__.py +5 -0
- dew/nn/autoencoders/api.py +154 -0
- dew/nn/autoencoders/dc_ae.py +513 -0
- dew/nn/autoencoders/flux2.py +100 -0
- dew/nn/autoencoders/kl.py +108 -0
- dew/nn/autoencoders/pretrained.py +94 -0
- dew/nn/autoencoders/qwen_image.py +389 -0
- dew/nn/autoencoders/rae.py +516 -0
- dew/nn/autoencoders/sd_vae.py +64 -0
- dew/nn/autoencoders/vae.py +434 -0
- dew/nn/autoencoders/wan.py +509 -0
- dew/nn/backbones/__init__.py +30 -0
- dew/nn/backbones/causal_transformer.py +1942 -0
- dew/nn/backbones/decoder_block.py +989 -0
- dew/nn/backbones/decoder_stack.py +537 -0
- dew/nn/backbones/dit.py +119 -0
- dew/nn/backbones/edm2.py +162 -0
- dew/nn/backbones/flux.py +191 -0
- dew/nn/backbones/flux2.py +199 -0
- dew/nn/backbones/jepa.py +220 -0
- dew/nn/backbones/joint.py +296 -0
- dew/nn/backbones/layer_plan.py +192 -0
- dew/nn/backbones/mmdit.py +303 -0
- dew/nn/backbones/qwen_image.py +261 -0
- dew/nn/backbones/sd3.py +165 -0
- dew/nn/backbones/ssm_dit.py +83 -0
- dew/nn/backbones/unet.py +135 -0
- dew/nn/backbones/unet3d.py +105 -0
- dew/nn/backbones/unet_condition.py +312 -0
- dew/nn/backbones/uvit.py +261 -0
- dew/nn/backbones/video_dit.py +63 -0
- dew/nn/backbones/wan.py +250 -0
- dew/nn/backbones/z_image.py +235 -0
- dew/nn/blocks.py +285 -0
- dew/nn/conv.py +251 -0
- dew/nn/deepseek_v4.py +957 -0
- dew/nn/diffusion_gemma.py +357 -0
- dew/nn/dit.py +666 -0
- dew/nn/dsa_kpool.py +450 -0
- dew/nn/dspark.py +190 -0
- dew/nn/engram.py +326 -0
- dew/nn/fake_quant.py +117 -0
- dew/nn/gemma3n.py +225 -0
- dew/nn/gemma4_moe.py +99 -0
- dew/nn/gpt_oss.py +122 -0
- dew/nn/hyper_connections.py +206 -0
- dew/nn/inputs.py +720 -0
- dew/nn/kda.py +283 -0
- dew/nn/kernels/__init__.py +16 -0
- dew/nn/kernels/decode_attention.py +159 -0
- dew/nn/kernels/delta_rule.py +114 -0
- dew/nn/kernels/generation.py +96 -0
- dew/nn/kernels/grouped_matmul.py +181 -0
- dew/nn/kernels/ragged_dot.py +444 -0
- dew/nn/kernels/ssd.py +350 -0
- dew/nn/kv_cache.py +569 -0
- dew/nn/linear.py +671 -0
- dew/nn/llama4.py +104 -0
- dew/nn/mixer_base.py +125 -0
- dew/nn/mixers/__init__.py +4 -0
- dew/nn/mixers/attention.py +980 -0
- dew/nn/mixers/gated_delta_net.py +50 -0
- dew/nn/mixers/mamba2.py +592 -0
- dew/nn/mixers/mlp.py +88 -0
- dew/nn/mla.py +783 -0
- dew/nn/mobilenet.py +431 -0
- dew/nn/moe.py +1150 -0
- dew/nn/mp.py +171 -0
- dew/nn/multimodal.py +494 -0
- dew/nn/precision.py +252 -0
- dew/nn/protocols.py +304 -0
- dew/nn/rope.py +384 -0
- dew/nn/safety.py +56 -0
- dew/nn/scan_orders.py +148 -0
- dew/nn/scatter.py +13 -0
- dew/nn/sharding.py +639 -0
- dew/nn/sparse_selection.py +140 -0
- dew/nn/ssm.py +276 -0
- dew/nn/text_encoders.py +858 -0
- dew/nn/vision/__init__.py +137 -0
- dew/nn/vision/common.py +118 -0
- dew/nn/vision/deepseek_v41.py +212 -0
- dew/nn/vision/gemma3n.py +235 -0
- dew/nn/vision/gemma4.py +394 -0
- dew/nn/vision/llama4.py +296 -0
- dew/nn/vision/qwen35.py +342 -0
- dew/nn/vision/siglip.py +258 -0
- dew/objectives/__init__.py +4 -0
- dew/objectives/base.py +1055 -0
- dew/objectives/diffusion/__init__.py +28 -0
- dew/objectives/diffusion/adversarial.py +330 -0
- dew/objectives/diffusion/alignment.py +192 -0
- dew/objectives/diffusion/block.py +439 -0
- dew/objectives/diffusion/config.py +511 -0
- dew/objectives/diffusion/consistency.py +446 -0
- dew/objectives/diffusion/end_to_end.py +272 -0
- dew/objectives/diffusion/few_step.py +317 -0
- dew/objectives/diffusion/guidance_distillation.py +134 -0
- dew/objectives/diffusion/masked.py +242 -0
- dew/objectives/diffusion/objective.py +867 -0
- dew/objectives/distillation.py +256 -0
- dew/objectives/jepa/__init__.py +16 -0
- dew/objectives/jepa/config.py +86 -0
- dew/objectives/jepa/masking.py +148 -0
- dew/objectives/jepa/objective.py +258 -0
- dew/objectives/jepa/probes.py +149 -0
- dew/objectives/lm/__init__.py +4 -0
- dew/objectives/lm/chunked.py +862 -0
- dew/objectives/lm/config.py +288 -0
- dew/objectives/lm/objective.py +1276 -0
- dew/objectives/rl/__init__.py +93 -0
- dew/objectives/rl/episodes.py +737 -0
- dew/objectives/rl/flow.py +442 -0
- dew/objectives/rl/grpo.py +363 -0
- dew/objectives/rl/journal.py +158 -0
- dew/objectives/rl/ppo.py +325 -0
- dew/objectives/rl/preference.py +117 -0
- dew/objectives/rl/records.py +70 -0
- dew/objectives/rl/rollout.py +154 -0
- dew/objectives/rl/scheduler.py +572 -0
- dew/objectives/rl/sessions.py +674 -0
- dew/objectives/rl/sources.py +372 -0
- dew/objectives/rl/verl.py +314 -0
- dew/objectives/supervised.py +138 -0
- dew/pool.py +123 -0
- dew/position.py +98 -0
- dew/py.typed +0 -0
- dew/records.py +138 -0
- dew/registry.py +1250 -0
- dew/rl/__init__.py +46 -0
- dew/rl/_sandbox_exec.py +28 -0
- dew/rl/advantage.py +170 -0
- dew/rl/sandbox.py +670 -0
- dew/rl/surrogate.py +363 -0
- dew/sampling/__init__.py +72 -0
- dew/sampling/decoding.py +968 -0
- dew/sampling/flow.py +216 -0
- dew/sampling/guidance.py +379 -0
- dew/sampling/guided.py +164 -0
- dew/sampling/pipelines.py +854 -0
- dew/sampling/sample.py +144 -0
- dew/sampling/solvers/__init__.py +25 -0
- dew/sampling/solvers/brownian.py +176 -0
- dew/sampling/solvers/common.py +144 -0
- dew/sampling/solvers/dpm.py +388 -0
- dew/sampling/solvers/gaussian.py +230 -0
- dew/sampling/solvers/sigma.py +300 -0
- dew/sampling/solvers/unipc.py +213 -0
- dew/sampling/strategies.py +899 -0
- dew/sampling/text.py +953 -0
- dew/sampling/vocabulary.py +188 -0
- dew/telemetry/__init__.py +0 -0
- dew/telemetry/devices.py +160 -0
- dew/telemetry/instrumentation.py +453 -0
- dew/telemetry/peaks.py +51 -0
- dew/telemetry/profile.py +341 -0
- dew/telemetry/records.py +149 -0
- dew/training/__init__.py +25 -0
- dew/training/display.py +540 -0
- dew/training/distributed.py +905 -0
- dew/training/evaluation.py +493 -0
- dew/training/execution.py +429 -0
- dew/training/host.py +255 -0
- dew/training/memory.py +341 -0
- dew/training/narrow.py +99 -0
- dew/training/optim.py +896 -0
- dew/training/posthoc.py +93 -0
- dew/training/quantization.py +870 -0
- dew/training/rungs.py +70 -0
- dew/training/runtime.py +323 -0
- dew/training/selection.py +50 -0
- dew/training/state.py +82 -0
- dew/training/tracker.py +661 -0
- dew/training/trainer.py +2266 -0
- dew/training/transaction.py +486 -0
- dewml-0.1.0.dist-info/METADATA +1594 -0
- dewml-0.1.0.dist-info/RECORD +335 -0
- dewml-0.1.0.dist-info/WHEEL +5 -0
- dewml-0.1.0.dist-info/entry_points.txt +2 -0
- dewml-0.1.0.dist-info/licenses/LICENSE +21 -0
- 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())
|