Dew is a JAX framework for training language models, image and video diffusion models, and JEPA encoders. It comes with common model architectures and training objectives, and one trainer handles optimization, sharding across devices, evaluation and checkpoints for all of them.
Dew started as a fork of FlaxDiff, the diffusion-only framework I have been building and refining for years. With FlaxDiff I trained Stable Diffusion-style text-to-image models from scratch on more than 400M text-image pairs, on 128 TPU v4 chips.
You can use the built-in architectures, load a supported Hugging Face checkpoint, or train your own Flax model. Variables and training state are ordinary JAX PyTrees, optimizers are Optax transformations, data loading uses Grain, and checkpoints use Orbax.
Dew is not at 1.0 yet, so APIs and checkpoint formats can still change. Models lists what is supported and the known limits.
Contents
- Getting started
- Features
- Models
- Training
- Diffusion and sampling
- Generating and serving
- Distributed training
- Data and configuration
- Installation
- Documentation and examples
- Contributing and acknowledgements
Getting started
This section trains a diffusion transformer on Oxford Flowers at 64×64 on an NVIDIA GPU, then samples a grid of images.
Install Dew and CUDA JAX in a virtual environment:
git clone https://github.com/AshishKumar4/dew.git
cd dew
uv venv --python 3.14
source .venv/bin/activate
uv pip install -e ".[tfds,cuda12]"
Prepare the dataset once, in a separate environment. The TFDS builder needs TensorFlow, but training reads the prepared files without it. TFDS 4.9.10 imports importlib_resources while it prepares a dataset but only declares that dependency for Python before 3.9, so the command installs it explicitly.
uv venv --python 3.13 .venv-data
uv pip install --python .venv-data/bin/python \
"tensorflow-datasets==4.9.10" "tensorflow==2.21.0" scipy importlib_resources
CUDA_VISIBLE_DEVICES="" .venv-data/bin/python - <<'PY'
from pathlib import Path
import tensorflow_datasets as tfds
builder = tfds.builder(
"oxford_flowers102",
data_dir=Path.home() / ".cache" / "dew" / "datasets",
)
builder.download_and_prepare(file_format="array_record")
print(builder.data_dir)
PY
Here is an abridged examples/train_flowers.py. The full file adds command-line options and the sampling step:
from pathlib import Path
import jax
import jax.numpy as jnp
import optax
from dew import Checkpoints, Field, InputSpec, Trainer
from dew.data import Loading, TFDSImages
from dew.diffusion.presets import EDM
from dew.objectives.diffusion import DiffusionObjective
from dew.nn.backbones import SimpleDiT
def train():
data_path = Path.home() / ".cache/dew/datasets/oxford_flowers102/2.1.1"
data = TFDSImages(
path=str(data_path),
split="train",
image_size=64,
val_batches=0,
loading=Loading(workers=2, threads=2, read_buffer=16, worker_buffer=2),
).load(batch=16)
model = SimpleDiT(
patch_size=4,
emb_features=128,
num_layers=4,
num_heads=4,
dtype=jnp.bfloat16,
attention_impl="auto",
)
objective = DiffusionObjective(
model,
EDM(regime="pixel"),
InputSpec(Field("image", (64, 64, 3))),
)
trainer = Trainer(
objective,
optax.adamw(2e-4),
key=jax.random.key(0),
checkpoints=Checkpoints("runs/flowers64/checkpoints"),
)
return trainer.fit(data, steps=1000, log_every=20, checkpoint_every=200)
if __name__ == "__main__":
state = train()
The __main__ guard lets Grain start its data-loading worker processes. Field describes one image, and the dataset yields batches of 16. The objective adds noise and builds the denoising targets. The trainer runs the optimizer and writes checkpoints, and state.averaged holds the EMA weights you sample from.
Run the full script. It also saves a sample grid to runs/flowers64/samples.png:
CUDA_VISIBLE_DEVICES=0 JAX_PLATFORMS=cuda python examples/train_flowers.py \
--data "$HOME/.cache/dew/datasets/oxford_flowers102/2.1.1" \
--steps 1000
Use --steps 20 for a short run, or increase --steps to train longer.
examples/train_diffusion.py adds pretrained CLIP text conditioning, and examples/train_flowers_tpu.py runs the same job across a TPU slice and scores the result. For an offline run with no dataset download, examples/readme_demo.py covers language modeling, resuming from a checkpoint, DPO and flow matching.
Change the training setup
Pass an Optax optimizer to Trainer. To use momentum SGD in the Flowers script, replace its optimizer argument:
optimizer = optax.sgd(learning_rate=1e-2, momentum=0.9)
trainer = Trainer(
objective,
optimizer,
key=jax.random.key(0),
checkpoints=Checkpoints("runs/flowers-sgd/checkpoints"),
)
Set ema_decay when you construct the objective. A value closer to 1 averages weights over more updates:
objective = DiffusionObjective(
model,
EDM(regime="pixel"),
InputSpec(Field("image", (64, 64, 3))),
ema_decay=0.999,
)
After training, sample from the averaged weights with state.averaged. Use state.variables instead to sample from the latest weights:
from dew import sample
from dew.sampling import Heun
process = objective.process
denoise = process.denoiser(model, state.averaged, conditions={})
images = sample(
denoise,
process.noise(jax.random.key(1), (8, 64, 64, 3)),
solver=Heun(),
steps=40,
key=jax.random.key(2),
)
Set ema_decay=None to train without an averaged copy, then sample with state.variables.
SimpleDiT(dtype=jnp.bfloat16) sets the dtype the model computes in. The
master weights and optimizer state stay in fp32, so updates still accumulate in
full precision.
For int8 quantization-aware training, install the quantization extra
(pip install "dewml[quantization]", which installs Qwix) and wrap the model before you construct the objective:
from dew.training import Quantization
model = Quantization(dtype="int8", patterns=(".*dit_block_.*",)).apply(model)
The pattern selects the DiT's transformer blocks for int8 and leaves the
patch-embedding convolution in its floating-point dtype. Master weights stay in
fp32. Change patterns to quantize other modules.
In code, you build a model from its class (from dew.nn.backbones import SimpleDiT, CausalTransformer). A saved run records the class by its import path, which a recipe imports to rebuild the run, and dew.registry.models holds short aliases such as simple_dit for the command line.
Features
| Area | What Dew provides |
|---|---|
| Language modeling | Autoregressive pretraining, packed documents, assistant-only SFT, DPO, and GRPO with callable rewards |
| Diffusion | Image and video denoising, rectified flow, latent diffusion, masked-token diffusion, classifier-free guidance, and interchangeable schedules and solvers |
| Representation learning | I-JEPA and V-JEPA with context and target encoders, predictors, block masking, and linear and kNN probes |
| Training systems | Data, FSDP, expert, tensor and sequence parallelism, layer scans, gradient accumulation, mixed precision, asynchronous checkpoints, and configurable EMA |
| Interoperability | Hugging Face configuration and weight translation for the supported families, safetensors export, CLIP and T5 conditioning, and VAE components |
| Evaluation | Perplexity, FID, CLIP score, PSNR, SSIM, representation diagnostics, generated previews, and Weights & Biases tracking |
DPO and GRPO train with the same Trainer as pretraining. dew.rl has PPO's advantage estimators and loss terms as separate functions, so you can build your own policy loop from them.
attention_impl selects the attention kernel: "reference", "xla", "cudnn", "tpu" (Pallas splash attention) or "auto". With "auto", Dew picks a kernel each time it traces an attention call, in this order:
- The reference path, if the call asks for arithmetic that no fused kernel does (a matmul precision above default, a softmax outside fp32, or a compute dtype different from the inputs'), has float64 inputs, or runs bf16 on a GPU older than sm80.
- cuDNN, on a supported GPU, if the call has no attention sinks.
- Splash attention, on a TPU, if the call qualifies: lengths of at least 512 tokens that it can tile, a mask it can describe, no additive bias, and whole sequences at the kernel.
- XLA otherwise.
The kernel never changes the parameter tree, so a checkpoint trained with one kernel loads with any other.
Models
Supported models lists every checkpoint
family Pretrained.load can read, by the model_type in its config.json,
and every architecture Dew can train from scratch. The site builds that page
from Dew's registries. The notes below cover training, inference and export
for each group.
Text decoders
Pretrained.load reads the checkpoint's config.json and builds a
CausalTransformer. Training uses LMObjective, generation uses
dew.sampling.generate or PretrainedDecoder.text_generation(), and
Pretrained.save writes config.json, model.safetensors and
generation_config.json back in the Hugging Face layout.
dtype sets the compute dtype and param_dtype sets the dtype parameters
are stored in. By default Pretrained.load keeps parameters in FP32. Pass
param_dtype="bfloat16" to halve weight memory without changing the compute
dtype. State that is not a parameter keeps its declared precision.
When exported, Kimi K2 keeps its own model type, vocabulary, RoPE settings and
routing widths. Small fixtures test loading, a Trainer update, export, and
reloading in the reference implementation; their
source record pins the released
configuration. Kimi K2.5 wraps the same decoder in a vision model. Dew loads and
trains the text decoder and writes the vision tower and projector tensors back
byte for byte, but it does not run the vision part.
For Qwen3-Next, GLM-5.3 and GLM-5.3-Flash, tiny fixtures built from the released
configurations match their transformers classes, including the prediction
layer, and export back in the source layout after a Trainer update. Dew reads
Mamba 2 from both the Hugging Face port (Mamba2ForCausalLM) and the original
mamba_ssm checkpoints such as state-spaces/mamba2-130m, and saves in the
port's layout. A decoder built on Llama's block that mixes sliding-window and
full attention layers exports as ministral, which transformers can load.
For Kimi K3, Dew loads the text decoder out of the vision wrapper. That covers
the KDA and NoPE MLA layers, Attention Residuals over blocks of layers, latent
routed experts with SiTU, and the routed experts' compressed-tensors MXFP4
weights. Export writes the experts back as the same packed pairs, re-encoding
trained experts with the library's own rule
(dew.interop.codecs.quantize_packed_mxfp4). The vision tower tensors are
written back unchanged. A tiny fixture built from the released remote code
tests parity, an update, export and greedy decoding, and the
source record lists the shape
of every released tensor.
Kimi Linear (KDA and NoPE MLA layers with DeepSeek's routed experts) computes what its released remote code computes, with one exception. The released gate adds the balancing bias to its scores in place, so the routing weights include the bias. Dew weights the chosen experts by the unbiased scores, as vLLM and K3's revision of the same file do. A tiny fixture from the released code, with that line patched, tests parity, an update, export and greedy decoding; the source record lists the shape of every released tensor.
Native multimodal models
For a multimodal checkpoint, Pretrained.load returns the model, the
checkpoint's own processor and the weights. The processor turns text and raw
media into ModelInputs, which LMObjective, Trainer and cached generation
accept as they are. Export writes the processor and tokenizer files next to
the weights.
The checkpoint's modality configuration decides which image, video and audio
inputs a model accepts. Processor.__call__ takes text, images, audio,
videos and video_metadata. Processor.chat applies the checkpoint's own
chat template, so template controls such as reasoning_effort and
preserve_thinking behave as the checkpoint defines them. The checkpoint's
processor also does the image, video and waveform preprocessing, and Dew
arranges its outputs row by row. DeepSeek-V4.1 ships without a processor, so
you build its pixel values and image positions yourself
(language models).
Qwen 3.8 uses the Qwen 3.5 model types. Qwen/Qwen3.8-27B loads as a
qwen3_5 conditional model with a dense hybrid decoder and image and video
inputs. The text-only Qwen/Qwen3.8-2.4T-A95B loads as qwen3_5_moe_text,
with normalized top-k routing and a sigmoid-gated shared expert.
tests/fixtures/hf/qwen38-source/source.json pins both revisions, and every
tensor name in their indexes maps to a Dew parameter, including the
multi-token prediction (MTP) layer they ship. The MTP layer shares the target
embedding and head, trains as an auxiliary loss, and decodes candidate steps
from its own cache. Dew refuses a checkpoint with more than one prediction
layer.
I did not download the released Qwen 3.8 weights. The tests run tiny fixtures with the released shapes, on CPU in float32, so I have not checked full-size memory use, bf16 parity, accelerator throughput, multi-host placement or a speculative accept/reject scheduler.
Block-diffusion decoders
Diffusion Gemma (diffusion_gemma) generates canvases and trains with text
and image-conditioned SFT.
BlockDiffusionObjective trains the canvas loss starting from the loaded
weights, and PretrainedBlockDecoder.block_generation() decodes canvases. Text
SFT follows Google's published recipe. Image-conditioned SFT goes through the
same ModelInputs, Dataset and Trainer. The images condition the clean
encoder, and their placeholder slots are not text targets. To decode, pass the
matching images= to BlockGeneration. The objective makes layer_scalar
trainable, so export has to use the objective's model:
replace(loaded, model=objective.model).save(directory, variables=state.variables).
Masked-diffusion decoders
LLaDA (llada) and Dream (dream, Dream) train with the MDLM loss,
starting from their released weights.
MaskedDiffusionObjective(loaded, MDLM(mask_id=...)(), seq_len) trains on the
MDLM negative ELBO starting from a loaded checkpoint.
loaded.save(directory, variables=state.variables) writes
the trained weights back with the source's own tensor names (OLMo-style names
for LLaDA, the Qwen 2 layout for Dream), next to the config they came with.
The transformers library has no class for either release, so the export tests compare
against the transformers model each one is built from, run with an all-visible
attention mask: LlamaForCausalLM for LLaDA and Qwen2ForCausalLM for Dream.
Pretrained diffusion and quantized checkpoints
Pretrained.load reads SD, SDXL, SD3, Flux and Qwen-Image 2.1 pipeline
directories. For SD and SDXL that includes img2img, inpainting and the SDXL
refiner.
The loader reads these quantized storage formats:
- DeepSeek's FP8 blocks (
weight_scale_inv); - DeepSeek-V4 and V4.1's
.scalestorage (FP8 layers, FP4 routed experts, engram tables); - GPT-OSS's MXFP4;
- compressed-tensors'
mxfp4-pack-quantized,nvfp4-pack-quantized(weights only),pack-quantized(int codes of 1 to 8 bits),float-quantized(FP8 by tensor, channel or block),int-quantizedandnaive-quantized; - AutoAWQ's 4-bit gemm packing;
- GPTQ at 2, 4 and 8 bits, act-order included.
It decodes each format to the same values as the release's own
dequantization. Pretrained.save writes trained weights back in the source's
format and scale dtype, with the encoding rule of the tool that wrote the
source.
AWQ and GPTQ weights are saved against the scales and zeros the source shipped. A change smaller than half a grid step is lost, so a lightly trained model mostly saves back as its source. Dew refuses to save a trained value outside the grid, because AutoAWQ's packing would spill it into the neighbouring codes and gptqmodel's would clamp it. Save a substantially trained model dense instead; the error message names the call. A compressed-tensors weight is saved the way the library's own compressor writes it: from the weight in the dtype Dew holds, against the source's scales, clamping a value past the code range as the library does. Activations that a checkpoint quantizes dynamically run in the model's dtype, and Dew refuses static activation scales.
The model-family tests use small fixtures with the source's shapes. I have not validated full-size checkpoints, accelerator performance or physical multi-host runs. Video inputs are tested on Gemma 4 and Qwen 3.5.
Diffusion and representation models
Dew trains its own UNets, DiTs, MMDiTs, a video DiT and the I-JEPA and V-JEPA encoders from scratch, for diffusion, flow matching and masked representation prediction; Supported models lists them by alias.
CLIP and T5 text encoders supply conditioning, and the VAE interfaces let you train in latent space.
Training
A training run combines a model, an objective, a dataset, and an optimizer:
| Object | Responsibility |
|---|---|
| Flax model | Forward computation and variable structure |
Objective |
Initialization, loss, evaluation outputs, and optional previews |
Dataset |
Training and validation iterator factories |
Trainer |
Differentiation, optimizer updates, placement, checkpoints, and logging |
TrainState |
Parameters, optimizer state, random key, progress, and EMA variables |
Language modeling
This decoder trains on TinyStories, a corpus of short stories in simple English, with the GPT-2 tokenizer. Download the 22 MB validation file of TinyStories V2 and tokenize it into the train.bin, val.bin and meta.json files that TokenWindows reads. dew tokenize holds out the first 1% of the tokens for validation.
hf download roneneldan/TinyStories TinyStoriesV2-GPT4-valid.txt \
--repo-type dataset --local-dir data
dew tokenize --input data/TinyStoriesV2-GPT4-valid.txt \
--out data/tinystories --tokenizer gpt2
Each row holds 257 token IDs: the model receives the first 256 and predicts the following 256.
import jax
import jax.numpy as jnp
import optax
from dew import Trainer
from dew.data import HFTokenizer, Loading, TokenWindows
from dew.nn.backbones import CausalTransformer
from dew.objectives.lm import LMObjective, Perplexity
from dew.sampling import Sampling, generate
tokenizer = HFTokenizer("gpt2")
data = TokenWindows(path="data/tinystories", seq_len=256,
loading=Loading(workers=0)).load(batch=32)
model = CausalTransformer(vocab_size=tokenizer.vocab_size, emb_features=256, num_layers=4,
num_heads=4, max_seq_len=256, dtype=jnp.bfloat16)
objective = LMObjective(model, seq_len=256)
lm_state = Trainer(objective, optax.adamw(1e-3), key=jax.random.key(0)).fit(
data, steps=2000, log_every=500, eval_every=1000, metrics=(Perplexity(),))
continuation = generate(model, lm_state.variables, [tokenizer.encode("Once upon a time")],
max_new_tokens=40, key=jax.random.key(1),
sampling=Sampling(temperature=0.0))
print(tokenizer.decode(continuation.tokens[0]))
On one Colab L4 GPU the run takes about three minutes. The training loss falls from 2.75 at step 500 to 1.95 at step 2,000, and validation perplexity reaches 8.7. One run's greedy continuation reads:
Once upon a time, there was a little girl named Lily. She had a big, red ball. Lily loved to play with her ball. One day, she saw a big box. She wanted to open it.
temperature=0 selects the highest-probability token. GPU reductions are not bitwise repeatable by default, so a second run can continue differently after the first sentence. Validation uses EMA weights, which lag the live parameters during a short run: at step 1,000 their perplexity is 29.3.
To pack whole documents into the windows instead, set TokenWindows(pack=True). It adds segment IDs and positions, and splits the stream at the EOS ID that dew tokenize --pack records. ChatMessages reads conversations from a parquet file, a JSONL file or a Hub dataset ID, renders them with the tokenizer's chat template, and records each token's role. Set LMObjective(loss_role=Role.ASSISTANT) to train only on assistant targets. See language models for checkpoint loading and text tokenization.
Supervised fine-tuning
Next, fine-tune the decoder above on one response to a prompt. Each token has a role: the prompt's tokens are Role.USER and the response's are Role.ASSISTANT. When you read conversation data, ChatMessages builds this role column from the chat template.
import itertools
import numpy as np
from dew import Dataset
from dew.data.chat import Role
prompt = tokenizer.encode("Tom had a red ball.")
response = tokenizer.encode(" He kicked it to his dog.")
row = np.array(prompt + response, dtype=np.int32)
roles = np.array([Role.USER] * len(prompt) + [Role.ASSISTANT] * len(response), dtype=np.int8)
sft_batch = {"text": np.tile(row, (8, 1)), "text_roles": np.tile(roles, (8, 1))}
sft_data = Dataset(
train=lambda partition: itertools.repeat(sft_batch),
val=None,
records=8,
batch=8,
)
sft_objective = LMObjective(
model,
seq_len=len(row) - 1,
variables=lm_state.variables,
loss_role=Role.ASSISTANT,
)
sft_state = Trainer(
sft_objective,
optax.adamw(1e-3),
key=jax.random.key(2),
).fit(sft_data, steps=20, log_every=10)
The loss counts only the assistant targets, after the next-token shift; the prompt tokens are still there as context. For conversation files, ChatMessages also keeps tool calls, tool responses and tool schemas.
Preference optimization
This continues from model and lm_state above, with a chosen and a rejected response to the same prompt. The masks limit the loss to the response tokens.
import json
from dew.data import PreferencePairs
from dew.objectives.rl import DPOObjective
rejected = tokenizer.encode(" He kicked kicked kicked it.")
pair = {"chosen": prompt + response, "rejected": prompt + rejected,
"chosen_mask": [0] * len(prompt) + [1] * len(response),
"rejected_mask": [0] * len(prompt) + [1] * len(rejected)}
pairs = PreferencePairs(records=(json.dumps(pair),) * 8, seq_len=16,
loading=Loading(workers=0, threads=1, read_buffer=2)).load(batch=8)
dpo = DPOObjective(model, seq_len=15, beta=0.1, variables=lm_state.variables)
dpo_state = Trainer(dpo, optax.adam(0.001), key=jax.random.key(2)).fit(
pairs, steps=10, log_every=5)
DPOObjective keeps the starting policy as a frozen reference and raises the likelihood of the chosen response relative to the rejected one. PreferencePairs.seq_len is the width of the whole ID row, and shorter pairs are padded to it. The objective scores one position fewer because of the next-token shift.
FlowGRPOObjective applies group-relative rewards to stochastic flow trajectories. FlowRollout samples groups of images, computes their rewards, and records the transition densities that the clipped policy objective uses. FlowGRPO has a complete example with an image reward.
For online reinforcement learning, SampledRollout generates groups of responses and calls a reward function on them. GRPOObjective then trains on their advantages, old log probabilities and response masks. recipes/chain.py connects SFT, DPO, and GRPO stages. The post-training guide covers reward callbacks and rollout settings.
Reinforcement learning with a reward function
This example keeps training the same decoder and rewards stories that mention a dog. A task verifier could replace the reward function. The prompt batch has the same numeric layout that Prompts produces, including the UTF-8 reward metadata.
from dew.objectives.rl import GRPOObjective, SampledRollout
prompt = tokenizer.encode("Once upon a time, there was a little")
prompt_batch = {
"prompt": np.tile(np.array(prompt, dtype=np.int32), (8, 1)),
"prompt_length": np.full(8, len(prompt), dtype=np.int32),
"data_source": np.tile(
np.frombuffer(b"tinystories", dtype=np.uint8).astype(np.int32),
(8, 1),
),
"ground_truth": np.tile(
np.frombuffer(b"dog", dtype=np.uint8).astype(np.int32),
(8, 1),
),
"extra_info": np.zeros((8, 0), dtype=np.int32),
}
def reward(data_source, completion, ground_truth, extra_info):
return float(ground_truth in completion)
rl_data = Dataset(
train=lambda partition: itertools.repeat(prompt_batch),
val=None,
records=8,
batch=8,
)
rl_objective = GRPOObjective(
model,
seq_len=len(prompt) + 7,
beta=0.01,
variables=lm_state.variables,
)
rollout = SampledRollout(
rl_objective,
reward=reward,
groups=4,
max_new_tokens=8,
sampling=Sampling(temperature=1.0, top_k=40),
decode=tokenizer.decode,
)
rl_state = Trainer(
rl_objective,
optax.adamw(1e-4),
key=jax.random.key(3),
rollout=rollout,
).fit(rl_data, steps=20, log_every=10)
Each prompt gets four responses, and SampledRollout.decode turns each one into the text the reward reads. The advantages come from how each response's reward compares with the others in its group. GRPO uses a clipped policy objective and an optional KL term against the reference; beta sets the KL coefficient. seq_len covers the prompt and the 8 response tokens, minus one for the next-token shift.
In the Colab run, 3 of 128 responses sampled from the policy before these 20 steps mention a dog, and all 128 sampled after them do.
Diffusion language models
Masked diffusion trains a bidirectional decoder to recover masked tokens. There is no next-token shift, so each input row holds exactly seq_len tokens. Windows of seq_len=127 hold 128 IDs each, which is what the 128-token objective needs. The mask token takes the first ID after GPT-2's vocabulary.
from dew.diffusion.discrete import MDLM
from dew.objectives.diffusion import MaskedDiffusionObjective
masked_data = TokenWindows(path="data/tinystories", seq_len=127,
loading=Loading(workers=0)).load(batch=64)
mask_id = tokenizer.vocab_size
process = MDLM(mask_id=mask_id)()
masked_model = CausalTransformer(
vocab_size=mask_id + 1,
emb_features=256,
num_layers=4,
num_heads=4,
max_seq_len=128,
causal=False,
dtype=jnp.bfloat16,
)
masked_objective = MaskedDiffusionObjective(masked_model, process, seq_len=128)
masked_state = Trainer(
masked_objective,
optax.adamw(1e-3),
key=jax.random.key(4),
).fit(masked_data, steps=4000, log_every=1000, eval_every=2000,
metrics=(Perplexity(),))
drawn = process.generate(masked_model, masked_state.averaged,
[tokenizer.encode("Once upon a time")], 48, key=jax.random.key(5))
print(tokenizer.decode(drawn.tokens[0]))
On the same L4 the 4,000 steps take about five minutes. The loss, MDLM's negative ELBO per token, falls from 3.48 at step 1,000 to 2.78 at step 4,000, and validation perplexity reaches 15.0. That perplexity is exp of the ELBO, an upper bound on the model's own, so it does not compare directly with the 8.7 above. process.generate unmasks the 48 tokens after the prompt in 64 reverse steps. One run's sample reads:
Once upon a time, in a small town, there lived a little girl named Lily. Mia loved whistle. She had an key telling to try to sleep. One day, she always to see her favorite mom, dad unate to have lots of.
LLaDA and Dream train the same way. Diffusion Gemma uses a different process, with canvases and self-conditioning.
JEPA representation learning
A JEPA encoder learns by predicting the representations of masked image patches. This example uses the Flowers data prepared in Getting started, with 8×8 patches and a smaller predictor.
from pathlib import Path
import jax
import optax
from dew import Field, Trainer
from dew.data import Loading, TFDSImages
from dew.objectives.jepa import JepaEncoder, JepaObjective, JepaPredictor, KnnProbe, MultiBlockMask
def train_jepa():
data = TFDSImages(
path=str(Path.home() / ".cache/dew/datasets/oxford_flowers102/2.1.1"),
split="train",
image_size=64,
val_batches=2,
loading=Loading(workers=0, threads=1, read_buffer=2),
).load(batch=16)
encoder = JepaEncoder(
patch_size=8,
emb_features=64,
num_layers=2,
num_heads=4,
)
predictor = JepaPredictor(
grid=(8, 8),
emb_features=64,
predictor_features=32,
num_layers=1,
num_heads=4,
)
objective = JepaObjective(
encoder,
predictor,
mask=MultiBlockMask.for_grid((8, 8), num_targets=1, scale=(0.25, 0.25)),
sample=Field("image", (64, 64, 3)),
momentum_steps=20,
)
return Trainer(
objective,
optax.adamw(1e-3),
key=jax.random.key(0),
).fit(
data,
steps=20,
log_every=10,
eval_every=20,
metrics=(KnnProbe(102),),
)
if __name__ == "__main__":
jepa_state = train_jepa()
The target encoder is an EMA of the context encoder. The loss compares predicted and target representations and never reconstructs pixels. The kNN metric is a quick check on half a batch; judge representation quality with a larger labeled evaluation. jepa_video_encoder applies the same objective to video clips.
Loading pretrained weights
You can load a supported checkpoint from a Hub repository or a local directory. This example downloads Qwen3-0.6B and its tokenizer (about 1.5 GB).
import jax
import jax.numpy as jnp
from transformers import AutoTokenizer
from dew.interop import PretrainedDecoder
from dew.sampling import Sampling, generate
checkpoint = "Qwen/Qwen3-0.6B"
tokenizer = AutoTokenizer.from_pretrained(checkpoint)
pretrained = PretrainedDecoder.load(
checkpoint,
dtype=jnp.bfloat16,
max_seq_len=512,
)
prompt = jnp.asarray(
[tokenizer.encode("Explain gradient accumulation in one paragraph.")],
dtype=jnp.int32,
)
result = generate(
pretrained.model,
pretrained.variables,
prompt,
max_new_tokens=128,
key=jax.random.key(0),
sampling=Sampling(temperature=0.8, top_k=40, eos_token_ids=tokenizer.eos_token_id),
)
print(tokenizer.decode(result.tokens[0], skip_special_tokens=True))
To train from those weights, pass the bundle in place of the model (LMObjective(pretrained, seq_len=512), or a post-training objective), and tokenize the training data with the checkpoint's own tokenizer. pretrained.adapt(LoRA(rank=8, modules=("q_proj", "v_proj")), key=0) adds a low-rank adapter first, so the objective trains only the adapter's factors (language models). Generating and serving shows how to sample from and export the result.
Composing a native decoder
CausalTransformer takes its layer pattern as configuration. This decoder
has three sliding-window attention layers followed by one global layer. It
trains on a Grain stream, then generates from the trained state:
import grain.python as grain
import jax
import jax.numpy as jnp
import numpy as np
import optax
from dew import Checkpoints, Dataset, Trainer
from dew.nn.backbones.causal_transformer import CausalTransformer
from dew.nn.backbones.layer_plan import LayerKind
from dew.objectives.lm import LMObjective
from dew.sampling import Sampling
model = CausalTransformer(
vocab_size=8, emb_features=64, num_layers=4,
num_heads=4, num_kv_heads=2, mlp_features=128, max_seq_len=32,
layer_types=("sliding_attention",) * 3 + ("full_attention",),
kinds={"sliding_attention": LayerKind(window=8)},
dtype=jnp.float32,
)
row = np.resize(np.array([1, 2, 3, 4], np.int32), 17)
batch = {"text": np.tile(row, (4, 1))}
stream = grain.MapDataset.source([batch]).repeat().to_iter_dataset()
data = Dataset(train=lambda partition: iter(stream), val=None, records=4, batch=4)
objective = LMObjective(model, seq_len=16)
checkpoints = Checkpoints("runs/custom-decoder")
state = Trainer(objective, optax.adamw(0.003), key=jax.random.key(0),
checkpoints=checkpoints).fit(data, steps=40, log_every=20,
checkpoint_every=40)
checkpoints.wait()
task = objective.pipeline(state)
result = task([[1, 2]], 8, key=jax.random.key(1), sampling=Sampling(temperature=0))
print(np.asarray(result.tokens))
layer_types names each layer's kind, and kinds says what that kind does:
sliding_attention attends to 8 keys including its own, and full_attention
attends to all of them. A Grain to_iter_dataset() iterator has
get_state and set_state, which checkpoint_every needs to record the data
position. A plain generator has neither, so Trainer raises an error if you
use one with checkpoint_every. objective.pipeline returns a
TextGeneration task that uses the state's weights.
Custom Flax models
An objective can train any Linen module. This example fits y = 2x + 1 with a single dense layer.
import flax.linen as nn
import jax
import jax.numpy as jnp
import numpy as np
import optax
from dew import Aux, Dataset, Field, InputSpec, Objective, Trainer
class Regression(Objective):
model = nn.Dense(1)
inputs = InputSpec(Field("x", (1,)))
def init(self, key, variables=None):
return self.model.init(key, jnp.ones((1, 1)))
def loss(self, variables, batch, step):
prediction = self.model.apply(variables, batch["x"])
loss = jnp.mean((prediction - batch["y"]) ** 2)
return loss, Aux(metrics={"mse": loss})
x = np.linspace(-1, 1, 32, dtype=np.float32).reshape(32, 1)
data = Dataset.from_records({"x": x, "y": 2 * x + 1}, batch=32)
objective = Regression()
state = Trainer(objective, optax.sgd(0.1), key=jax.random.key(0)).fit(
data, steps=100, log_every=25)
print(objective.model.apply(state.variables, jnp.array([[0.0], [1.0]])))
The predictions approach 1 and 3. init creates the variables, loss returns a differentiable scalar, and Aux holds metrics and any updates to mutable variables. You can pass a custom objective straight to Trainer, and a configuration names it by its import path.
The objective guide also covers BatchNorm state and EMA selection.
Evaluation and checkpoints
Pass eval_every and metrics to fit to score the validation data, and set preview=True if you also want generated previews. Perplexity is computed from the token losses of the whole validation pass, and FID accumulates statistics over the pass before it computes the distance.
Checkpoints saves the training state and the data position with Orbax. To continue a run, rebuild it with the same checkpoint directory. fit(steps=1200) is a total: a run restored at step 1000 trains to step 1200, not 2200. The recipe configuration is saved separately; RunConfig.save writes it to run.json.
See checkpointing and resume for restore requirements, local checkpoints, and current recovery limitations.
Which call continues a run depends on the artifact you kept:
| What you have | What it continues into | Call |
|---|---|---|
| A native checkpoint directory | The same run: optimizer state, the root key, and the data position | Trainer(..., checkpoints=Checkpoints(directory)), then fit |
A TrainState in memory |
Generation from the weights you just trained | objective.pipeline(state) |
A saved run: run.json beside its checkpoints |
Generation, with the model rebuilt from the record | dew.pipeline(run_directory) |
| A source checkpoint directory or Hub repository | Generation, or training from those weights | dew.pipeline(source), or Pretrained.load(source) for the variables |
| Trained variables another runtime has to read | The source format, without Dew | Pretrained.save(directory, variables=state.variables) |
| A saved run another runtime has to read | The same layout, from the run alone | PretrainedDecoder.from_run(run_directory).save(destination), or dew export <run> <dest> |
Pretrained.save writes the weights, the config it derives, and the
tokenizer or processor files. It does not include optimizer state or data
position, so keep the native checkpoint if you want to resume training.
examples/sft_gemma4.py trains a run and exports
it with PretrainedDecoder.from_run(...).save(...), the last row of the table.
To reproduce a run bit for bit on CUDA you need deterministic GPU reductions.
The flag --xla_gpu_deterministic_ops=true turns them on, and
TrainerConfig.xla_flags appends it to XLA_FLAGS. Check the flag with your
attention backend: under JAX 0.11.1, repeated cuDNN backward calls fail with it
set, while the XLA attention path passed the recorded bitwise checks.
tests/test_training_qualification.py resumes a killed fine-tune on the XLA path.
Standalone evaluation and local reports
Evaluation.run scores trained variables without an optimizer and returns metric values and optional previews. LocalTracker writes the scalar history, artifacts and plots to disk, with no W&B account or install. Install dewml[plots] for Matplotlib output.
import itertools
import jax
import numpy as np
import optax
from dew import Dataset, Evaluation, LocalTracker, Trainer
from dew.nn.backbones import CausalTransformer
from dew.objectives.lm import LMObjective, Perplexity, Samples
from dew.sampling import Sampling
row = np.resize(np.array([1, 2, 3, 4], dtype=np.int32), 17)
batch = {"text": np.tile(row, (8, 1))}
data = Dataset(
train=lambda partition: itertools.repeat(batch),
val=lambda partition: iter([batch]),
records=8,
batch=8,
)
model = CausalTransformer(
vocab_size=8,
emb_features=32,
num_layers=1,
num_heads=2,
mlp_features=64,
max_seq_len=32,
)
objective = LMObjective(
model,
seq_len=16,
samples=Samples([1, 2], 8, sampling=Sampling(temperature=0)),
)
with LocalTracker("runs/lm-report", plots=True) as tracker:
state = Trainer(
objective,
optax.adam(0.01),
key=jax.random.key(0),
tracker=tracker,
).fit(data, steps=40, log_every=10)
result = Evaluation.run(
objective,
state.variables,
data.val,
metrics=(Perplexity(),),
key=jax.random.key(1),
step=int(state.step),
preview=True,
)
tracker.log(result.scalars, step=result.step)
for preview in result.previews:
tracker.artifact(preview, step=result.step)
print(result.scores)
This run reports a perplexity of about 1.002 and saves the training-loss curve, the scalar journal and the generated text under runs/lm-report. The tracker draws the plots once, when it closes. Use plots=False to record only scalars and artifacts, or call tracker.plot() yourself. examples/evaluate_and_serve.py scores a finished run the same way and adds an lm-eval-harness suite, image metrics, and a served-model comparison.
Trackers sends the same reports to several backends. To switch backends, change the constructor. Install dewml[wandb], dewml[mlflow] or dewml[tensorboard] for the backend you want:
from dew import LocalTracker, TensorBoardTracker, Trackers
tracker = Trackers(
LocalTracker("runs/experiment/tracking"),
TensorBoardTracker("runs/experiment/events"),
)
Use it in the same with block and Trainer call as above. WandbTracker(project="dew-experiments", offline=True) and MLflowTracker("dew-experiments", uri="sqlite:///runs/mlflow.db") go in the same place. A custom backend implements log, artifact and close. Run configuration, progress, checkpoint requests, profiler windows, sweep trials and failures reach the tracker as typed records.
Profiling training and inference
Install dewml[profile], or run uv pip install -e '.[profile]' in this checkout. Then wrap the work you want to profile in dew.Profiler, either as a context manager or with start and stop. Both forms record the same JAX/XProf capture:
import dew
with dew.Profiler("profiles/run"):
state = trainer.fit(data, steps=1000)
prof = dew.Profiler("profiles/run")
prof.start()
try:
state = trainer.fit(data, steps=1000)
finally:
prof.stop()
Each capture goes in a new directory, so restarting a profiler keeps the earlier results. Without a path, the first start creates a temporary directory that is not deleted, available as prof.directory. A capture keeps the native XPlane traces, the HLO files that exist, and XProf's overview, input, kernel, memory and other supported reports. Its manifest records the backend, package versions, capture options and which reports are available. A counter the backend does not provide is recorded as missing, not as zero. The profile extra installs XProf's viewer, and each manifest stores the command that opens its capture under view_command, for example xprof --logdir=profiles/run/capture-<id>.
A capture leaves JAX's Python tracer off. That tracer records every Python and C call and slows Python-heavy host work several times over, so the host time in its traces is not time the run would spend without it. To trace differently, pass Profiler an options= value. To start a trace yourself with the same settings, use jax.profiler.start_trace(directory, profiler_options=capture_options()), with capture_options from dew.telemetry.profile.
To trace a window of training, pass Trainer a ProfileWindow with the trace directory, the number of steps to trace and the warmup steps to run first. The loop starts tracing after the warm-up, stops after the requested steps, and reports the window to the tracker as a ProfileWindow record. Use this schedule or an outer dew.Profiler, but not both.
Sweeping a hyperparameter
RunConfig.sweep trains one trial for each point of a search space with the ordinary RunConfig.train. It keeps a JSON ledger so the sweep can resume, and reports each trial to the tracker you pass it:
from dew import Evaluation, LocalTracker
from dew.config import ModelConfig, OptimConfig, RunConfig, TrainerConfig
from dew.config.sweep import GridSearch
from dew.data import TokenWindows
config = RunConfig(
# The run records the model it trains, and the batches above stand in
# for the dataset this names.
model=ModelConfig.from_model(objective.model),
data=TokenWindows(seq_len=16),
optim=OptimConfig(optimizer="adam"),
trainer=TrainerConfig(name="lm-rate", checkpoint_dir="runs/sweep", steps=40, batch_size=8,
eval_every=None, checkpoint_every=None),
)
def trial(run: RunConfig) -> float:
"""Train one point and score it: the perplexity its own run ends on."""
state = run.train(objective, data, name=run.trainer.name or "lm-rate")
return float(Evaluation.run(objective, state.variables, data.val, metrics=(Perplexity(),),
key=jax.random.key(1), step=int(state.step)).scores["val/perplexity"])
with LocalTracker("runs/sweep/tracking") as tracker:
trials = config.sweep({"optim.learning_rate": [0.01, 0.003]}, train=trial, trials=2,
ledger="runs/sweep/ledger.json", tracker=tracker, search=GridSearch())
best = min(trials, key=lambda trial: trial.value)
print(best.overrides, round(best.value, 4))
This prints {'optim.learning_rate': 0.01} 1.0024; the slower rate reaches 1.015. Each trial is a real run under runs/sweep/lm-rate/trial-<index>, with its own run.json, checkpoints and tracking journal. A finished trial is written to the ledger before it is reported, so calling sweep again continues an interrupted sweep without retraining the finished trials. RandomSearch and GridSearch are built in; OptunaSearch needs dewml[hpo].
Diffusion and sampling
A Process combines a noise schedule, a prediction transform and a loss weighting. Presets build the common combinations:
| Component | Options |
|---|---|
| Presets | EDM, Karras, Cosine, Flow, Sqrt; MDLM for masked-token diffusion |
| Prediction transforms | Noise, clean-sample, velocity, flow, Karras preconditioning |
| Weighting | Schedule weighting, P2-related weighting, Min-SNR |
| Solvers | DDPM/DDIM, Euler/Heun/RK4, DPM-Solver and DPM-Solver++, DEIS, UniPC, PNDM, LMS, KDPM2, EDM-DPM, LCM, TCD, DPM-Solver SDE |
| Guidance | Classifier-free guidance with an optional interval and rescaling |
| Conditions | InputSpec/Condition, CLIP, T5, labels or custom encoders |
Training and sampling can use different schedules; EDM, for example, trains on a log-normal distribution of noise levels and samples on a Karras grid. sample runs the solver under jax.lax.scan, and you can change the solver without retraining. MultiStepDPM integrates in sigma space and keeps the previous denoiser outputs to raise the order of each step.
TextToImage runs text encoding, denoising and, for latent models, decoding. See diffusion for text conditioning, latent models, and sampling.
For SD and SDXL checkpoints, dew.pipeline(source) rebuilds the source's own
scheduler rather than choosing a solver by name. It supports the source's
clipping, thresholding and timestep spacing, and the Karras, exponential and
beta grids where the matching scheduler has them. An unsupported combination
raises an error; Dew does not fall back to different defaults.
For a text-conditioned image task, pass
guidance=CFG(scale=7.5, interval=(0.0, 1.0), rescale=0.7)
after importing CFG from dew.sampling. Rescaling mixes the guided output
with a version matched to the conditional output's standard deviation.
rescale=0.0 leaves it unchanged, and guidance=None turns classifier-free guidance off.
The tests that compare these schedulers with the source use tiny synthetic
trajectories; they are not quality benchmarks of the released models.
Generating and serving
A task in dew.inference pairs a model with its weights for generation.
TextGeneration decodes tokens, BlockGeneration decodes Diffusion Gemma
canvases, MaskedGeneration samples a whole response from a masked-diffusion
decoder with Dew's MDLM sampler, and TextToImage denoises images.
MaskedGeneration does not implement LLaDA's or Dream's own remasking recipes.
A task is immutable and keeps the weights it was built with; bind returns a
new task with other weights.
Drawing from a trained language model
LMObjective.policy returns a TextGeneration over the parameters you pass
it. The GRPO rollout samples with the same task. This continues the decoder
trained in language modeling:
from dew.inference import TextGeneration
prompt = tokenizer.encode("Once upon a time")
task = objective.policy(lm_state.variables, Sampling(temperature=0.0))
drawn = task([prompt], 8, key=jax.random.key(1))
print(tokenizer.decode(drawn.tokens[0]), np.asarray(drawn.lengths))
same = TextGeneration(model, lm_state.variables, sampling=Sampling(temperature=0.0))
print(np.array_equal(np.asarray(same([prompt], 8, key=jax.random.key(1)).tokens),
np.asarray(drawn.tokens)))
This prints Once upon a time, there was a little girl named Lily [8] and
True, because the constructor and policy build the same task. lengths
counts the 8 generated tokens and leaves out the 4 prompt tokens in the same
row. Generation also returns terminated, plus the behavior_log_probs and
raw_log_probs that a policy ratio needs.
A task built by PretrainedDecoder.text_generation() includes the checkpoint's
processor, so it accepts strings and decode returns text. Without a
processor the task takes token rows or ModelInputs.
Drawing from a trained diffusion run
TextToImage.from_run loads a run directory. It reads the run.json that a
DiffusionRunConfig wrote and the weights of the latest checkpoint, with the
EMA copy merged over the live parameters, so you don't repeat the model
configuration at generation time.
from pathlib import Path
import jax
import jax.numpy as jnp
import optax
from dew import Checkpoints, Trainer
from dew.config import ModelConfig, TrainerConfig
from dew.data import Loading, TFDSImages
from dew.diffusion import presets
from dew.inference import TextToImage
from dew.nn.backbones import SimpleDiT
from dew.objectives.diffusion import DiffusionRunConfig
from dew.sampling import Heun
run = Path("runs/flowers-run")
config = DiffusionRunConfig(
model=ModelConfig.from_model(SimpleDiT(patch_size=4, emb_features=128, num_layers=4,
num_heads=4, dtype=jnp.bfloat16)),
data=TFDSImages(
path=str(Path.home() / ".cache/dew/datasets/oxford_flowers102/2.1.1"),
image_size=64,
val_batches=0,
loading=Loading(workers=0, threads=1, read_buffer=2),
),
trainer=TrainerConfig(checkpoint_dir=str(run), batch_size=16, steps=20, keep=1),
preset=presets.EDM(regime="pixel"),
text=None,
)
def main():
objective = config.build()
checkpoints = Checkpoints(str(run), keep=1)
state = Trainer(objective, optax.adamw(2e-4), key=jax.random.key(0),
checkpoints=checkpoints).fit(
config.data.load(batch=16, tokenize=objective.inputs.tokenize),
steps=20, log_every=20, checkpoint_every=20)
checkpoints.wait()
config.save(str(run))
print(sorted(path.name for path in run.iterdir()))
task = TextToImage.from_run(str(run))
images = task(["a flower", "another flower"], steps=20, solver=Heun(),
key=jax.random.key(1))
print(images.host().images.shape, int(state.updates))
if __name__ == "__main__":
main()
The run directory then holds ['20', 'run.json'] and the task draws
(2, 64, 64, 3) images clipped to [-1, 1]. text=None trains an
unconditional model, so the prompt list only sets how many images to draw; a
run with a TextCondition encodes the prompts with the encoder it names.
config.data.load takes tokenize=objective.inputs.tokenize because the
objective's conditions read the dataset's captions. Loading(workers=0) keeps
that caption reader in the training process, so it shuts down with the run.
Exporting a decoder and serving it
PretrainedDecoder.from_model(model, variables, tokenizer=...) wraps a model
you trained in Dew so you can export it. save(directory) writes the weights
and config in the Hugging Face layout and saves the run's tokenizer files next
to them. Another runtime can load that directory directly.
Tokenize the corpus with the tokenizer you will export, so the token IDs match
the exported vocabulary. Here that is tiny-tools, a small byte-level BPE
tokenizer committed for the tests, and the corpus is the TinyStories file from
Language modeling:
dew tokenize \
--input data/TinyStoriesV2-GPT4-valid.txt \
--out runs/tokens \
--tokenizer tests/fixtures/tokenizers/tiny-tools
from pathlib import Path
import jax
import jax.numpy as jnp
import numpy as np
import optax
from dew import Trainer
from dew.data import HFTokenizer, Loading, TokenCorpus, TokenWindows
from dew.interop import PretrainedDecoder
from dew.nn.backbones import CausalTransformer
from dew.objectives.lm import LMObjective
from dew.sampling import Sampling
tokens = Path("runs/tokens")
export = Path("runs/dew-decoder")
corpus = TokenCorpus.read(tokens)
tokenizer = HFTokenizer(corpus.tokenizer)
data = TokenWindows(path=str(tokens), seq_len=128,
loading=Loading(workers=0, threads=1, read_buffer=2)
).load(batch=16)
model = CausalTransformer(vocab_size=corpus.vocab_size,
emb_features=128, num_layers=4, num_heads=4, num_kv_heads=2,
mlp_features=256, max_seq_len=128, dtype=jnp.float32,
qk_norm=False, tie_embeddings=False)
state = Trainer(LMObjective(model, seq_len=128),
optax.adamw(3e-3), key=jax.random.key(0)).fit(
data, steps=400, log_every=200)
PretrainedDecoder.from_model(model, state.variables, tokenizer=tokenizer).save(export)
print(sorted(path.name for path in export.iterdir()))
task = PretrainedDecoder.load(str(export), dtype=jnp.float32).text_generation()
drawn = task("The trainer", 12, key=jax.random.key(1),
sampling=Sampling(temperature=0.0))
print(task.decode(drawn))
The export directory holds the weights, the config that both runtimes read, and the tokenizer's files:
runs/dew-decoder/
├── chat_template.jinja
├── config.json
├── generation_config.json
├── model.safetensors
├── tokenizer.json
└── tokenizer_config.json
With qk_norm=False and tie_embeddings=False the exporter writes a llama
config, an architecture llama.cpp can convert. The default decoder exports as
qwen3, which Ollama 0.32.9 rejects with unsupported architecture "Qwen3ForCausalLM". ollama create reads the export through the running
daemon, so start ollama serve first, then convert the directory with a
Modelfile:
cat > Modelfile <<'EOF'
FROM runs/dew-decoder
TEMPLATE "{{ .Prompt }}"
PARAMETER num_gpu 0
PARAMETER num_ctx 128
EOF
ollama create dew-decoder -f Modelfile
num_gpu 0 keeps the runner on the CPU. TEMPLATE "{{ .Prompt }}" passes
the prompt through unchanged, so you can compare the served output with the
local one.
These steps ran on Ollama 0.32.9 on Linux x86-64. On the same machine, Ollama
0.34.3's ollama create stops at MLX runtime is not available.
OllamaCompletion and OpenAICompletion wrap the vendors' SDK clients,
which you create yourself. Pass a Sampling to ask for the same sampling
settings that local generation uses. The client translates it into backend
options and turns off the daemon's own repetition penalties and other
truncations.
import ollama
from dew.inference import OllamaCompletion
from dew.sampling import Sampling
client = OllamaCompletion("dew-decoder", ollama.Client(host="http://127.0.0.1:11434"))
served = client("The trainer", 12, sampling=Sampling(temperature=0.0), key=0, raw=True)
print(served.texts, served.finish_reasons)
With temperature=0.0, the daemon reproduces Dew's greedy output token for
token, so served.texts[0] == task.decode(drawn)[0]. Completion also has
token_counts, a usage record, and the raw SDK responses under
responses. For vLLM, serve the same directory and pass provider="vllm",
which enables the sampling controls vLLM accepts beyond the OpenAI schema.
provider="sglang" does the same for SGLang:
vllm serve runs/dew-decoder --served-model-name dew-decoder
import openai
from dew.inference import OpenAICompletion
client = OpenAICompletion(
"dew-decoder",
openai.OpenAI(base_url="http://127.0.0.1:8000/v1", api_key="none"),
provider="vllm",
)
Without provider="vllm" or provider="sglang", the client raises an error for top_k, min_p and
eos_id instead of dropping them, because generic OpenAI endpoints do not
accept them.
Distributed training
MeshSpec describes how the devices are arranged, and Layout maps model dimensions onto that mesh. For example, MeshSpec(fsdp=4) splits eligible parameters and optimizer state over four devices, and the remaining devices form the data-parallel axis. Trainer works the same way on one device or a mesh.
The mesh also has expert, tensor, sequence and stage axes. Sequence-parallel attention chooses how to exchange data on each call. It uses Ulysses all-to-alls, which trade sequence rows for heads so that no device holds a whole key or value, when the heads and lengths divide evenly and the all-to-all moves fewer bytes, as it does for causal attention. Otherwise it gathers keys and values. The GPipe stage axis splits execution into stages, but parameters and optimizer state stay replicated across stages, so stages save activation memory and not parameter memory. MeshSpec(fsdp=8, replicas=2) is hybrid sharding over two nodes: FSDP inside each node and replicas across them.
Compute dtypes are configurable, and the attention kernel depends on the hardware: cuDNN for supported shapes on NVIDIA GPUs, a Pallas kernel on TPU, and XLA otherwise. Qwix provides optional int8 and fp8 compute, and MuonClip adds per-head QK clipping to Muon. Loading quantized weights is separate from quantized training.
Set remat on a decoder to a policy name such as "full", "minimal" or "save_qkv_proj" to recompute block activations in the backward pass. A policy such as "save_qkv_proj" keeps the projections it names. Offloaded policies such as "minimal_offloaded" keep those residuals in pinned host memory, and Layout(host=("opt_state", "ema")) keeps the optimizer state and the EMA copy there between steps. Both save device memory at the cost of extra compute or transfers, and both work with layer scanning.
Start with distributed training and the TPU guide. Benchmarks and performance notes record workload sizes, hardware, memory, and timing.
Multiple hosts
Every host runs the same script. The example below uses two hosts with two GPUs each and shards the model state over all four devices. Both hosts must see the same token files and checkpoint directory. Training on several nodes covers Slurm, hybrid sharding and long sequences.
Prepare byte-token data from your corpus and place it on shared storage:
dew tokenize \
--input corpus.txt \
--out /shared/tokens \
--tokenizer byte \
--val-fraction 0.01
Save this as train_multihost.py. prepare_process joins the process pool, so call it before you create any device array.
import os
import jax
import optax
def main():
from dew.training.runtime import prepare_process
prepare_process(multi_host=True)
try:
import jax.numpy as jnp
from dew import Checkpoints, MeshSpec, Trainer
from dew.data import Loading, TokenWindows
from dew.nn.backbones import CausalTransformer
from dew.objectives.lm import LMObjective
data = TokenWindows(
path=os.environ["DEW_TOKEN_DIR"],
seq_len=128,
loading=Loading(workers=0, threads=1, read_buffer=2),
).load(batch=16)
model = CausalTransformer(
vocab_size=256,
emb_features=128,
num_layers=2,
num_heads=4,
mlp_features=256,
max_seq_len=128,
dtype=jnp.bfloat16,
)
trainer = Trainer(
LMObjective(model, seq_len=128),
optax.adamw(3e-4),
key=jax.random.key(0),
mesh=MeshSpec(fsdp=4),
checkpoints=Checkpoints(os.environ["DEW_CHECKPOINT_DIR"]),
)
state = trainer.fit(
data,
steps=int(os.environ.get("DEW_STEPS", "1000")),
log_every=20,
checkpoint_every=200,
)
print(f"Process {jax.process_index()}: {int(state.updates)} updates")
finally:
jax.distributed.shutdown()
if __name__ == "__main__":
main()
Launch it from the first host. dew launch uses ssh to start one process per GPU on each host (four in all) and gives each process its GPU, the coordinator address, the process count and its rank. The remote shell does not read a login profile, so give the interpreter's absolute path:
dew launch --hosts 10.0.0.1,10.0.0.2 \
--env DEW_TOKEN_DIR=/shared/tokens --env DEW_CHECKPOINT_DIR=/shared/runs/lm \
-- /opt/dew/.venv/bin/python train_multihost.py
batch=16 is the global batch, so each process reads four rows. Without --hosts, the same command runs on this machine's GPUs. Inside a Slurm allocation it starts srun, and with --tpu NAME it runs on every worker of a Cloud TPU. Training on several nodes covers each case. To rehearse the launch on a machine without GPUs, run dew launch --processes-per-host 2 --env JAX_PLATFORMS=cpu --env XLA_FLAGS=--xla_force_host_platform_device_count=2 --env DEW_STEPS=20 ... with local directories. Two processes with two simulated devices each make up the same MeshSpec(fsdp=4), and each process prints 20 updates.
Generating text with Gemma 4 on one GPU
dew.pipeline loads a published checkpoint and returns a task you can call. Set JAX_PLATFORMS=cuda before starting Python. I have not run this released Gemma 4 checkpoint on the 4080, so check that it fits in memory before you try it.
import jax.numpy as jnp
import dew
chat = dew.pipeline("google/gemma-4-E2B-it", dtype=jnp.bfloat16)
result = chat(
["Explain gradient accumulation in one paragraph.",
"Name three uses of a JEPA encoder."],
128,
key=0,
)
for text in result.text:
print(text)
The task includes the checkpoint's processor, so it accepts strings and result.text holds the decoded continuations. The second positional argument is the token budget; if the checkpoint's generation_config.json declares one, that is the default. n=4 draws four continuations per prompt. The prompts above go straight to the tokenizer. To apply the checkpoint's chat template, call task.processor.chat(messages) and pass the ModelInputs it returns to the same call.
E2B wraps a multimodal model, and the task also accepts images= when the checkpoint declares a vision tower. The dtype here sets the compute dtype, not how weights are stored: the Gemma 4 loader keeps FP32 parameters. Loading also needs host buffers, device temporaries and the KV cache, so this example does not show that E2B fits on a 16 GB GPU.
Generating text with a large decoder on a TPU slice
Every process runs the same script with its own prompts, and result.host() returns only that process's real rows. I have not run this on a real TPU slice. The pipeline loads the checkpoint before it shards it, so having enough device memory in total does not mean loading will succeed.
Save this as generate_tpu.py:
import os
import jax
import jax.numpy as jnp
def main():
jax.distributed.initialize() # Cloud TPU workers discover the coordinator
try:
import dew
from dew.training import Layout, MeshSpec
task = dew.pipeline(
os.environ.get("DEW_MODEL", "google/gemma-4-31B-it"),
mesh=MeshSpec(fsdp=jax.device_count()),
layout=Layout(min_shard=2**16),
dtype=jnp.bfloat16,
)
rank = jax.process_index()
result = task([f"Process {rank}: write one sentence about tensors."], 64, key=0)
for text in result.host().text:
print(rank, text)
finally:
jax.distributed.shutdown()
if __name__ == "__main__":
main()
MeshSpec(fsdp=jax.device_count()) shards eligible weights over the slice; small leaves, and leaves that don't divide evenly, may stay replicated. On Cloud TPU, jax.distributed.initialize() can find the coordinator by itself. On a cluster you manage yourself, pass the coordinator address, process count and process ID, as in the training example. Authenticate each worker through its environment or credential store, and keep tokens out of launch arguments. Save the script on every worker, then preview the launch on every worker (see Cloud TPUs):
dew launch --tpu dew-16 --zone us-central2-b --dry-run \
--env DEW_MODEL=google/gemma-4-31B-it -- python generate_tpu.py
Each worker needs access to the checkpoint and tokenizer files. A shared download cache still leaves every process with its own loading buffers. When you estimate memory, count the stored weight dtype, replicated leaves, loading peaks and the KV cache; checkpoint bytes divided by the device count is too low. Row counts, tokenized shapes and execution settings must match across processes. Use the same seed on every rank, because Dew derives each global row's key from it. Rehearse with a small checkpoint on a CPU process pool before a real TPU run.
Data and configuration
Dew's dataset specifications build batches with Grain. The token loaders read fixed windows or packed documents, and the chat, preference and prompt loaders read post-training data. Image and video data can come from local files, Hugging Face datasets, TFDS, ArrayRecord shards or URL streams.
Loading sets the number of workers, read threads and buffers. Dataset also accepts your own iterator factories, as in the fine-tuning examples above. Data loading covers transforms, deterministic randomness, batching and who owns the iterators.
The recipes turn dataclass configurations into command-line options with tyro. ModelConfig, OptimConfig and TrainerConfig hold the model, optimizer and run settings, and task-specific configurations add diffusion or language-model options. --help lists the options:
python recipes/lm/train.py --help
python recipes/diffusion/train.py --help
python recipes/jepa/train.py --help
The recipe guide walks through preparing a text corpus and training on it.
Installation
I recommend Python 3.14. Dew requires Python 3.12 or later, and CI tests both versions. Install the release from PyPI with uv, with the extra for your hardware. Each extra installs the accelerator build of the JAX version Dew requires:
| Hardware | Install |
|---|---|
| CPU | uv pip install dewml |
| NVIDIA GPU | uv pip install "dewml[cuda12]" (or cuda13 for CUDA 13 drivers) |
| Google TPU | uv pip install "dewml[tpu]" |
To work on Dew itself, install it from a clone instead:
git clone https://github.com/AshishKumar4/dew.git
cd dew
uv venv --python 3.14
source .venv/bin/activate
uv pip install -e ".[cuda12]"
To install the main branch without cloning it, run uv pip install "dewml[cuda13] @ git+https://github.com/AshishKumar4/dew". The extras install the accelerator build of jax 0.11.2, the version Dew requires. A later -U "jax[...]" would replace it with a release Dew isn't tested on. A process pool across GPUs keeps its compilation cache only with a patched jax 0.11.2 (jax-ml/jax#40940); the installation guide explains how to install it. See the JAX installation guide for driver requirements.
The optional extras are av, cuda12, cuda13, diffusers, eval-harness, gguf, guided, hpo, inference-clients, interop, metrics, mlflow, mutation, plots, profile, quantization, serve, streaming, tensorboard, test, tfds, torch, torchax, tpu, vision, wan and wandb. interop reads and writes safetensors, vision provides the host image processors that the multimodal checkpoints call, and inference-clients installs the Ollama and OpenAI SDKs used in the serving section. The sections above name the extra each feature needs. The installation guide covers development dependencies and dataset preparation.
Documentation and examples
- Getting started: training a custom model.
- Custom objectives: loss functions and mutable model state.
- Language models: tokens, checkpoints, and generation.
- Post-training: SFT, DPO, GRPO, and rewards.
- Diffusion: image, video, conditioning, and latents.
- Representation learning: JEPA encoders and predictors.
- API reference: constructors, arguments, and state contracts.
- Examples and recipes: complete programs to adapt.
Each of the five scripts below runs a whole job, from data to scored weights. By default they use settings for real hardware; --smoke swaps those for the repository's tiny fixtures, a few steps and one CPU device. End-to-end examples gives both command lines for each.
examples/train_flowers_tpu.py: a text-to-image DiT trained on Oxford Flowers across a TPU slice, then sampled and scored with FID and CLIPScore.examples/sft_diffusion_gemma.py: LoRA SFT of DiffusionGemma with the base weights held in host memory, publishing a PEFT adapter directory.examples/sft_gemma4.py: full-weight SFT of a Gemma 4 decoder on a Hub chat dataset, exported to the Hugging Face layout.examples/train_rlvr.py: GRPO with verifiable rewards, where each completion is a program run against hidden tests in a sandbox fleet, and rollouts come from Dew's own server or a vLLM or SGLang server one update ahead of training.examples/evaluate_and_serve.py: perplexity, an lm-eval-harness suite, image metrics, and a served-model comparison over one finished run.
Contributing and acknowledgements
Questions, bug reports, and contributions are welcome. Read CONTRIBUTING.md before submitting changes. For numerical issues, include the model configuration, dependency versions, dtype, hardware, and a small reproduction.
Dew grew out of FlaxDiff. This project is partially supported by Google TPU Research Cloud. I would like to thank the Google Cloud TPU team for providing resources for the larger text-conditional experiments.
Dew builds on JAX, Flax, Optax, Grain, Orbax, tyro, and Weights & Biases. References and attribution lists the papers and upstream implementations used by its models and algorithms.
Dew is licensed under MIT. Adapted components, model weights, and datasets retain their applicable notices and licenses. If you use Dew in research, cite the repository and the papers for the models and methods you use.
Metadata
Release files for dewml 0.1.0
For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.
Source distribution (sdist)
| File | Size | Uploaded | |
|---|---|---|---|
| dewml-0.1.0.tar.gz | 1.6 MB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| dewml-0.1.0-py3-none-any.whl | Python 3 | none | any | Details |
Total release size: 3.3 MB
Release files / dewml-0.1.0.tar.gz
| Download URL | dewml-0.1.0.tar.gz |
|---|---|
| Size | 1.6 MB |
| Tags | Source |
|
SHA-256 checksum How to use checksums |
ce69f12bbd27873b275802919294c1978fc7c0d3843fc7eb56bac9891a06f509
|
|
BLAKE2b-256 checksum How to use checksums |
5e454730c60c06ee492d7f742a859a190a408b02e8f707655c653e57b151289d
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
Yes |
| Uploaded via |
twine/7.0.0 CPython/3.13.14
|
Provenance
Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.
PyPI Publish Attestation
PyPI verified that this artifact, at this checksum, originated from the publisher listed below.
Signed by GitHub Actions, verified by PyPI on Oct 9, 2026.
Transparency logRelease files / dewml-0.1.0-py3-none-any.whl
| Download URL | dewml-0.1.0-py3-none-any.whl |
|---|---|
| Size | 1.7 MB |
| Tags | Python 3 |
|
SHA-256 checksum How to use checksums |
352f9011835441a8ed8bb62d8526a77be07f2b7659b37eb1fe453226bd16290e
|
|
BLAKE2b-256 checksum How to use checksums |
e9b475588a2ce6a373d7b23c044a51a56d0f59659802bc60ade2f48789dd366a
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
Yes |
| Uploaded via |
twine/7.0.0 CPython/3.13.14
|
Provenance
Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.
PyPI Publish Attestation
PyPI verified that this artifact, at this checksum, originated from the publisher listed below.
Signed by GitHub Actions, verified by PyPI on Oct 9, 2026.
Transparency log