Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion src/maxtext/configs/base.yml
Original file line number Diff line number Diff line change
Expand Up @@ -806,6 +806,11 @@ olmo_apply_ngram_filter: true # mask instances with repetitive n-grams (OLMo-cor
# Training loop
steps: 150_001 # If set to -1 then will inherit value from learning_rate_schedule_steps
log_period: 100 # The frequency of Tensorboard flush, gcs metrics writing, and managed profiler metrics updating.
training_objective: 'causal_lm' # Supported objectives: causal_lm, block_diffusion
block_diffusion_mask_id: -1 # Tokenizer mask-token id; required for training_objective='block_diffusion'
block_diffusion_min_noise: 0.001 # Minimum per-block corruption probability for block-diffusion training
block_diffusion_logit_alignment: 'same_position' # Supported alignments: same_position, shifted
block_diffusion_canvas_policy: 'all_masked' # Supported canvases: all_masked, seed_and_mask

jax_distributed_initialization_timeout: 300 # This is the default timeout in https://github.com/jax-ml/jax/blob/main/jax/_src/distributed.py
# Note there are two separate initializations - the jax coordination service (aka jax.distributed.initialize) and the backend (e.g. PjRT), the timeout above refers
Expand Down Expand Up @@ -1352,4 +1357,3 @@ elastic_backup_kind: "snapshot"
elastic_timeout_seconds: 300
elastic_max_retries: 10
elastic_min_slice_count: -1

57 changes: 57 additions & 0 deletions src/maxtext/configs/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -1646,6 +1646,28 @@ class Distillation(BaseModel):
class TrainingLoop(BaseModel):
"""Configuration for the main training loop, evaluation, and reproducibility."""

training_objective: Literal["causal_lm", "block_diffusion"] = Field(
"causal_lm",
description="The token-prediction objective used to prepare targets and compute loss.",
)
block_diffusion_mask_id: int = Field(
-1,
description="The tokenizer mask-token id required by the block-diffusion training objective.",
)
block_diffusion_min_noise: float = Field(
1.0e-3,
gt=0.0,
le=1.0,
description="The minimum corruption probability sampled independently for each block.",
)
block_diffusion_logit_alignment: Literal["same_position", "shifted"] = Field(
"same_position",
description="How model logits align to clean target-token positions.",
)
block_diffusion_canvas_policy: Literal["all_masked", "seed_and_mask"] = Field(
"all_masked",
description="Whether every block is fully maskable or begins with a clean anchor token.",
)
steps: int = Field(
150_001,
ge=-1,
Expand Down Expand Up @@ -3521,6 +3543,41 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de
"Block-diffusion attention with attention='autoselected' or attention='flash' requires hardware='tpu'; "
"use attention='dot_product' on other hardware."
)
if self.training_objective == "block_diffusion":
if self.attention_type != AttentionType.BLOCK_DIFFUSION.value:
raise ValueError("`training_objective='block_diffusion'` requires `attention_type='block_diffusion'`.")
if self.block_diffusion_mask_id < 0 or self.block_diffusion_mask_id >= self.vocab_size:
raise ValueError(
f"`block_diffusion_mask_id` ({self.block_diffusion_mask_id}) must satisfy "
f"0 <= block_diffusion_mask_id < vocab_size ({self.vocab_size})."
)
# Block-diffusion attention validation above rejects packing first.
if self.packing: # pragma: no cover
raise ValueError("`training_objective='block_diffusion'` requires `packing=False`.")
if self.mtp_num_layers > 0:
raise ValueError("`training_objective='block_diffusion'` is not compatible with MTP.")
if self.num_vocab_tiling > 1:
raise ValueError("`training_objective='block_diffusion'` is not compatible with vocabulary tiling.")
if self.dataset_type != "hf":
raise ValueError("`training_objective='block_diffusion'` currently requires `dataset_type='hf'`.")
if self.use_dpo:
raise ValueError("`training_objective='block_diffusion'` is not compatible with DPO.")
if self.use_sft:
raise ValueError("`training_objective='block_diffusion'` currently supports pre-training only.")
if self.use_multimodal or self.use_audio:
raise ValueError("`training_objective='block_diffusion'` currently supports text-only training.")
valid_model_contracts = {
("same_position", "all_masked"),
("shifted", "seed_and_mask"),
}
model_contract = (self.block_diffusion_logit_alignment, self.block_diffusion_canvas_policy)
if model_contract not in valid_model_contracts:
raise ValueError(
"Block-diffusion training supports only `same_position/all_masked` or `shifted/seed_and_mask`; "
f"received `{model_contract[0]}/{model_contract[1]}`."
)
if self.block_diffusion_canvas_policy == "seed_and_mask" and self.causal_block_size < 2:
raise ValueError("`block_diffusion_canvas_policy='seed_and_mask'` requires `causal_block_size >= 2`.")
if self.quantize_kvcache and not self.kv_quant_axis:
raise ValueError("`kv_quant_axis` cannot be empty when quantize_kvcache is True.")
if (
Expand Down
66 changes: 63 additions & 3 deletions src/maxtext/input_pipeline/hf_data_processing.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,57 @@ def _get_pad_id(tokenizer):
return pad_id


def _get_training_objective_transform(
config: ml_collections.ConfigDict,
*,
shift: bool,
use_dpo: bool,
use_sft: bool,
packing: bool,
pad_id: int,
bos_token_id: int | None,
) -> input_pipeline_utils.ShiftData | input_pipeline_utils.BlockDiffusionCorruption | None:
"""Selects target preparation for causal or block-diffusion pre-training.

Args:
config: Training configuration containing the objective-specific settings.
shift: Whether causal language-model targets should be shifted by one token.
use_dpo: Whether the pipeline is preparing direct-preference data.
use_sft: Whether the pipeline is preparing supervised fine-tuning data.
packing: Whether multiple examples are packed into each sequence.
pad_id: Token ID used to pad causal language-model examples.
bos_token_id: Beginning-of-sequence token ID, or None when unavailable.

Returns:
The objective-specific Grain transform, or None when target shifting is disabled.

Raises:
ValueError: If the objective is unsupported or block diffusion is combined with
an incompatible post-training or packing mode.
"""
objective = getattr(config, "training_objective", "causal_lm")
if objective == "block_diffusion":
if use_sft:
raise ValueError("This block-diffusion integration currently supports pre-training only.")
if use_dpo:
raise ValueError("Block-diffusion pre-training is not compatible with DPO.")
if packing:
raise ValueError("Block-diffusion pre-training requires packing=False.")
return input_pipeline_utils.BlockDiffusionCorruption(
block_size=config.causal_block_size,
mask_id=config.block_diffusion_mask_id,
min_noise=config.block_diffusion_min_noise,
logit_alignment=config.block_diffusion_logit_alignment,
canvas_policy=config.block_diffusion_canvas_policy,
axis=1,
)
if objective != "causal_lm":
raise ValueError(f"Unsupported training objective: {objective}")
if shift and not use_dpo:
return input_pipeline_utils.ShiftData(ignored_ids=[pad_id, bos_token_id], axis=1)
return None


def vision_sft_preprocessing_pipeline(
dataset,
config,
Expand Down Expand Up @@ -354,11 +405,20 @@ def preprocessing_pipeline(
max_prompt_length = config.dpo.max_prompt_length
operations.append(dpo_utils.DPODataFormatting(pad_id, max_target_length, data_column_names, max_prompt_length))
else:
operations.append(input_pipeline_utils.PadOrTrimToMaxLength(max_target_length, pad_id))
operations.append(input_pipeline_utils.PadOrTrimToMaxLength(max_target_length, pad_id, config=config))
operations.append(grain.Batch(batch_size=batch_size, drop_remainder=drop_remainder))

if shift and not use_dpo:
operations.append(input_pipeline_utils.ShiftData(ignored_ids=[pad_id, tokenizer.bos_token_id], axis=1))
target_transform = _get_training_objective_transform(
config,
shift=shift,
use_dpo=use_dpo,
use_sft=use_sft,
packing=packing,
pad_id=pad_id,
bos_token_id=tokenizer.bos_token_id,
)
if target_transform is not None:
operations.append(target_transform)

# Since HuggingFace IterableDataset does not support access through index
# Indexes generated by dummy_index_sampler is not used.
Expand Down
66 changes: 63 additions & 3 deletions src/maxtext/input_pipeline/input_pipeline_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@
import numpy as np
from grain._src.python.dataset.sources.tfrecord_dataset import _TFRecordReader, _TFRecordDatasetIterator # pylint: disable=protected-access
from grain.experimental import TFRecordIterDataset
from maxtext.diffusion.block_diffusion import corruption as block_diffusion_corruption
from maxtext.input_pipeline.protos import example_pb2
from maxtext.input_pipeline import tokenizer
from maxtext.multimodal import processor as mm_processor
Expand Down Expand Up @@ -849,14 +850,22 @@ def map(
) -> dict[str, np.ndarray | mm_utils.PreprocessorOutput]:
"""map to each element"""
data_columns = list(element.keys())
preserve_pad_valued_tokens = (
self.config is not None and getattr(self.config, "training_objective", "causal_lm") == "block_diffusion"
)
for data_column in data_columns:
if data_column != "images":
if isinstance(element[data_column], mm_utils.PreprocessorOutput):
raise TypeError("Only 'images' column can be of type PreprocessorOutput.")

element[f"{data_column}_segmentation"] = (
element[data_column] != self.pad_id # pyrefly: ignore[unsupported-operation]
) # pyrefly: ignore[unsupported-operation]
if preserve_pad_valued_tokens:
element[f"{data_column}_segmentation"] = np.ones(
element[data_column].shape[0], dtype=np.int32 # pyrefly: ignore[missing-attribute]
)
else:
element[f"{data_column}_segmentation"] = (
element[data_column] != self.pad_id # pyrefly: ignore[unsupported-operation]
) # pyrefly: ignore[unsupported-operation]
# pyrefly: ignore[missing-attribute]
element[f"{data_column}_segmentation"] = element[
f"{data_column}_segmentation"
Expand All @@ -878,6 +887,8 @@ def map(

element["images"] = self._pad_image_and_mask(element["images"]) # pyrefly: ignore[bad-argument-type]

elif preserve_pad_valued_tokens and key.endswith(("_segmentation", "_position")):
element[key] = self._pad_text(element[key], self.max_length, 0) # pyrefly: ignore[bad-argument-type]
elif "true_length" not in key:
element[key] = self._pad_text(element[key], self.max_length, self.pad_id) # pyrefly: ignore[bad-argument-type]
return element
Expand Down Expand Up @@ -1001,6 +1012,55 @@ def map(self, element):
return shift_and_refine(element, ignored_ids=self.ignored_ids, axis=self.axis)


@dataclasses.dataclass
class BlockDiffusionCorruption(grain.RandomMapTransform):
"""Adapts block-diffusion corruption to the Grain batch contract."""

def __init__(
self,
block_size: int,
mask_id: int,
min_noise: float = 1.0e-3,
logit_alignment: str = "same_position",
canvas_policy: str = "all_masked",
axis: int = 1,
):
self.block_size = block_size
self.mask_id = mask_id
self.min_noise = min_noise
self.logit_alignment = logit_alignment
self.canvas_policy = canvas_policy
self.axis = axis

def random_map(self, element, rng: np.random.Generator):
"""Corrupts inputs while preserving clean targets and input metadata."""
inputs = np.asarray(element["inputs"])
targets = np.asarray(element["targets"])
targets_segmentation = np.asarray(element["targets_segmentation"])
if inputs.shape != targets.shape or inputs.shape != targets_segmentation.shape:
raise ValueError(
"inputs, targets, and targets_segmentation must have identical shapes, got "
f"{inputs.shape}, {targets.shape}, and {targets_segmentation.shape}"
)
result = block_diffusion_corruption.corrupt_tokens(
inputs,
targets_segmentation != 0,
rng,
block_size=self.block_size,
mask_id=self.mask_id,
min_noise=self.min_noise,
logit_alignment=self.logit_alignment,
canvas_policy=self.canvas_policy,
axis=self.axis,
)
output = dict(element)
output["inputs"] = result.inputs
output["targets"] = targets
output["corruption_mask"] = result.corruption_mask.astype(targets_segmentation.dtype)
output["targets_loss_mask"] = result.targets_loss_mask.astype(targets_segmentation.dtype)
return output


@dataclasses.dataclass
class ComputeQwen3OmniPositions(grain.MapTransform):
"""Computes 3D position IDs for Qwen3-Omni multimodal sequences.
Expand Down
57 changes: 50 additions & 7 deletions src/maxtext/trainers/pre_train/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,7 @@
from flax.nnx import variablelib

from maxtext.configs import pyconfig
from maxtext.diffusion.block_diffusion import target_alignment as block_diffusion_target_alignment
from maxtext.utils.globals import EPS
from maxtext.utils import elastic_utils
# Placeholder: internal
Expand Down Expand Up @@ -107,13 +108,32 @@ def loss_fn(model, config, data, dropout_rng, params, sparsity_state=None, is_tr
loss: average loss
aux: a dictionary including intermediate_outputs, xent_sum, and total_weights
"""
is_block_diffusion = getattr(config, "training_objective", "causal_lm") == "block_diffusion"
if getattr(config, "attention_type", "global") == "block_diffusion" and not is_block_diffusion:
raise ValueError(
"Block-diffusion attention requires target-aligned block-diffusion losses; "
"causal next-token labels would leak within a bidirectional block."
)
if is_block_diffusion:
required_masks = {"corruption_mask", "targets_loss_mask"}
missing_masks = required_masks - data.keys()
if missing_masks:
raise ValueError(f"Block-diffusion loss requires explicit batch masks; missing {sorted(missing_masks)}")
target_shape = data["targets"].shape
for mask_name in required_masks:
if data[mask_name].shape != target_shape:
raise ValueError(f"{mask_name} must match targets shape; got {data[mask_name].shape} and {target_shape}")

# decimate proportion of data when per_device_batch_size<1
if is_train:
for k, v in data.items():
data[k] = v[: config.micro_batch_size_to_train_on, :]
else:
for k, v in data.items():
data[k] = v[: config.micro_batch_size_to_eval_on, :]
if is_block_diffusion:
targets_loss_mask = (data["targets_loss_mask"] != 0) & (data["targets_segmentation"] != 0)
target_positions = data.get("targets_position", data["inputs_position"])
mutable_collections = ["intermediates"]
if config.mtp_num_layers > 0 and is_train:
# The single model.apply call now triggers the entire chain if MTP is enabled:
Expand Down Expand Up @@ -165,6 +185,13 @@ def loss_fn(model, config, data, dropout_rng, params, sparsity_state=None, is_tr
hidden_states = maxtext_utils.get_nested_value(intermediate_outputs, hidden_state_key)[0]
xent_sum, total_z_loss = vocab_tiling_linen_loss(hidden_states, data, config, model, params, is_train)
else:
if is_block_diffusion:
logits = block_diffusion_target_alignment.align_logits_to_targets(
logits,
config.block_diffusion_logit_alignment,
target_positions,
data["targets_segmentation"] != 0,
)
one_hot_targets = jax.nn.one_hot(data["targets"], config.vocab_size)
xent, z_loss = max_utils.cross_entropy_with_logits(logits, one_hot_targets, z_loss=config.z_loss_multiplier)

Expand All @@ -183,9 +210,12 @@ def loss_fn(model, config, data, dropout_rng, params, sparsity_state=None, is_tr
debug_sharding=config.debug_sharding,
)

# Mask out paddings at the end of each example.
xent = xent * (data["targets_segmentation"] != 0)
z_loss = z_loss * (data["targets_segmentation"] != 0)
if is_block_diffusion:
xent = xent * targets_loss_mask
z_loss = z_loss * targets_loss_mask
else:
xent = xent * (data["targets_segmentation"] != 0)
z_loss = z_loss * (data["targets_segmentation"] != 0)

xent_sum = jnp.sum(xent)
total_z_loss = jnp.sum(z_loss)
Expand Down Expand Up @@ -228,6 +258,13 @@ def loss_fn(model, config, data, dropout_rng, params, sparsity_state=None, is_tr
hidden_states = maxtext_utils.get_nested_value(intermediate_outputs, hidden_state_key)[0]
xent_sum, total_z_loss = vocab_tiling_nnx_loss(model, hidden_states, data, config, is_train)
else:
if is_block_diffusion:
logits = block_diffusion_target_alignment.align_logits_to_targets(
logits,
config.block_diffusion_logit_alignment,
target_positions,
data["targets_segmentation"] != 0,
)
one_hot_targets = jax.nn.one_hot(data["targets"], config.vocab_size)
xent, z_loss = max_utils.cross_entropy_with_logits(logits, one_hot_targets, z_loss=config.z_loss_multiplier)

Expand All @@ -246,14 +283,20 @@ def loss_fn(model, config, data, dropout_rng, params, sparsity_state=None, is_tr
debug_sharding=config.debug_sharding,
)

# Mask out paddings at the end of each example.
xent = xent * (data["targets_segmentation"] != 0)
z_loss = z_loss * (data["targets_segmentation"] != 0)
if is_block_diffusion:
xent = xent * targets_loss_mask
z_loss = z_loss * targets_loss_mask
else:
xent = xent * (data["targets_segmentation"] != 0)
z_loss = z_loss * (data["targets_segmentation"] != 0)

xent_sum = jnp.sum(xent)
total_z_loss = jnp.sum(z_loss)

total_weights = jnp.sum(data["targets_segmentation"] != 0)
if is_block_diffusion:
total_weights = jnp.sum(targets_loss_mask)
else:
total_weights = jnp.sum(data["targets_segmentation"] != 0)
# If gradient accumulation is enabled, we don't need to divide xent_sum
# by total_weights and then multiply the computed gradient by total_weights,
# since it's equivalent to computing the gradient from xent_sum.
Expand Down
Loading
Loading