Skip to content
Merged
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
3 changes: 2 additions & 1 deletion src/maxtext/configs/base.yml
Original file line number Diff line number Diff line change
Expand Up @@ -356,6 +356,7 @@ moe_mlpwo: 'remat'
query_proj: 'remat'
key_proj: 'remat'
value_proj: 'remat'
kv_proj: 'remat'
qkv_proj: 'remat'
out_proj: 'remat'
query_wa_proj: 'remat'
Expand Down Expand Up @@ -1133,7 +1134,7 @@ enable_prefix_caching: false
prefix_caching_hbm_byte: 10_000_000_000 # 10 GB
prefix_caching_dram_byte: 100_000_000_000 # 100 GB

# This is a temporary flag that will be removed soon after the fix lands in TE
# This is a temporary flag that will be removed soon after the fix lands in Transformer Engine
enable_padding_causal_mask: true

# Llama4-specific
Expand Down
88 changes: 64 additions & 24 deletions src/maxtext/configs/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -518,7 +518,10 @@ class ModelArchitecture(BaseModel):
True,
description="Whether to apply scale on query and key normalizations (default True).",
)
v_norm_with_scale: bool = Field(True, description="Whether to apply scale on value normalization (default True).")
v_norm_with_scale: bool = Field(
True,
description="Whether to apply scale on value normalization (default True).",
)


class MTP(BaseModel):
Expand Down Expand Up @@ -624,7 +627,7 @@ class Attention(BaseModel):
),
)
ragged_block_size: int = Field(256, description="Block size for ragged attention.")
enable_padding_causal_mask: bool = Field(True, description="Temporary flag for TE padding.")
enable_padding_causal_mask: bool = Field(True, description="Temporary flag for Transformer Engine padding.")
use_tokamax_splash: bool = Field(False, description="Whether to use tokamax splash attention.")
use_jax_splash: bool = Field(False, description="Whether to use jax splash attention.")
force_q_layout: bool = Field(False, description="Force the Q layout")
Expand Down Expand Up @@ -656,7 +659,10 @@ class CompressedAttention(BaseModel):

o_lora_rank: NonNegativeInt = Field(0, description="Output LoRA rank for Compressed Attention.")
o_groups: NonNegativeInt = Field(0, description="Output groups for Compressed Attention.")
compress_ratios: list[int] = Field(default_factory=list, description="Per-layer compression ratios (0, 4, 128, etc).")
compress_ratios: list[int] = Field(
default_factory=list,
description="Per-layer compression ratios (0, 4, 128, etc).",
)
compressed_rope_max_timescale: int = Field(
160000, description="If positive, used for Compressed Sparse/Heavy Attention."
)
Expand Down Expand Up @@ -757,14 +763,18 @@ class MoEGeneral(BaseModel):
num_experts: PositiveInt = Field(1, description="The total number of experts in each MoE layer.")
num_experts_per_tok: PositiveInt = Field(1, description="The number of experts to route each token to.")
capacity_factor: float = Field(-1.0, description="Expert capacity factor. If < 0, no token dropping.")
ragged_buffer_factor: float = Field(-1.0, description="Ragged buffer factor. If < 0, ragged buffer is worst case size.")
ragged_buffer_factor: float = Field(
-1.0,
description="Ragged buffer factor. If < 0, ragged buffer is worst case size.",
)
moe_expert_input_dim: int = Field(
-1,
description="Dimension of tokens entering the MoE layer. If < 0, defaults to emb_dim.",
)
base_moe_mlp_dim: int = Field(-1, description="Intermediate dimension at MoE layer.")
padded_base_moe_mlp_dim: Optional[int] = Field(
None, description="Padded intermediate dimension at MoE layer for efficient GMM_v2 kernel execution."
None,
description="Padded intermediate dimension at MoE layer for efficient GMM_v2 kernel execution.",
)
load_balance_loss_weight: NonNegativeFloat = Field(0.0, description="Weight for the load balancing auxiliary loss.")
use_custom_sort_vjp: bool = Field(
Expand All @@ -783,7 +793,8 @@ class MoEGeneral(BaseModel):
),
)
use_ragged_sort: bool = Field(
False, description="Whether to use ragged kernel for sorting, improve performance when EP is enabled."
False,
description="Whether to use ragged kernel for sorting, improve performance when EP is enabled.",
)
use_gather_mosaic_kernel: bool = Field(
False,
Expand Down Expand Up @@ -919,7 +930,8 @@ class DeepSeekMoE(BaseModel):
)
n_routing_groups: int = Field(-1, description="Number of groups for routing, disabled by default.")
first_num_hash_layers: int = Field(
0, description="Number of hash routing layers, used in DeepSeek V4 (0 means disabled)."
0,
description="Number of hash routing layers, used in DeepSeek V4 (0 means disabled).",
)
topk_routing_group: int = Field(-1, description="Number of top groups to route inputs to.")
use_batch_split_schedule: bool = Field(
Expand Down Expand Up @@ -993,7 +1005,8 @@ class HardwareAndMesh(BaseModel):
)
custom_mesh: str = Field("", description="Available options: ['hybrid_ring_64x4', 'hybrid_ring_32x8']")
custom_mesh_and_rule: CustomRule = Field(
CustomRule.DEFAULT, description="Customized mesh and logical rules for granularity."
CustomRule.DEFAULT,
description="Customized mesh and logical rules for granularity.",
)
allow_split_physical_axes: bool = Field(False, description="Allow splitting physical axes for device mesh creation.")
enable_nnx: bool = Field(True, description="Whether to use NNX for model definition.")
Expand All @@ -1002,7 +1015,8 @@ class HardwareAndMesh(BaseModel):
pure_nnx_decoder: bool = Field(True, description="Whether to enable pure NNX decoder.")
pure_nnx: bool = Field(True, description="Whether to enable pure NNX mode.")
remove_size_one_mesh_axis_from_type: bool = Field(
True, description="Whether to remove size one mesh axis from type through jax.config."
True,
description="Whether to remove size one mesh axis from type through jax.config.",
)


Expand All @@ -1028,7 +1042,10 @@ class LayoutAndSharding(BaseModel):
"with auto sharding, megablox kernel, and EP / FSDP parallelisms.",
)
shard_optimizer_over_data: bool = Field(False, description="Enable ZeRO-1 optimizer sharding over the data axis.")
internal_compile: bool = Field(False, description="Use internal_compile to bypass open-source topology mappings.")
internal_compile: bool = Field(
False,
description="Use internal_compile to bypass open-source topology mappings.",
)
internal_compile_num_devices: int = Field(-1, description="Number of devices when using internal_compile.")
compile_xla_flags: str = Field("", description="Compiler options for compilation only.")

Expand Down Expand Up @@ -1073,7 +1090,8 @@ class PipelineParallelism(BaseModel):
"""Configuration for pipeline parallelism."""

pipeline_fsdp_ag_per_repeat: bool = Field(
False, description="Enable weight prefetching for circular pipeline parallelism."
False,
description="Enable weight prefetching for circular pipeline parallelism.",
)
num_layers_per_pipeline_stage: int = Field(1, description="Number of layers to place on each pipeline stage.")
num_pipeline_repeats: int = Field(
Expand Down Expand Up @@ -1139,6 +1157,7 @@ class RematAndOffload(BaseModel):
query_proj: RematLocation = Field(RematLocation.REMAT, description="Remat policy for the query projection.")
key_proj: RematLocation = Field(RematLocation.REMAT, description="Remat policy for the key projection.")
value_proj: RematLocation = Field(RematLocation.REMAT, description="Remat policy for the value projection.")
kv_proj: RematLocation = Field(RematLocation.REMAT, description="Remat policy for the unified KV projection.")
query_wa_proj: RematLocation = Field(
RematLocation.REMAT,
description="Remat policy for the MLA query weighted attention projection.",
Expand Down Expand Up @@ -1321,7 +1340,10 @@ class OlmoGrainDataset(BaseModel):
``data_shuffle_seed``); only OLMo-specific fields are listed here.
"""

olmo_index_path: PathStr = Field("", description="Path or gs:// URI to the JSON index from build_olmo_npy_index.py.")
olmo_index_path: PathStr = Field(
"",
description="Path or gs:// URI to the JSON index from build_olmo_npy_index.py.",
)
olmo_path_remap_from: PathStr = Field(
"",
description="If set, rewrite index file paths starting with this prefix to olmo_path_remap_to.",
Expand Down Expand Up @@ -1428,19 +1450,24 @@ class Distillation(BaseModel):
distill_layer_indices: None | list = Field(None, description="Feature indices for feature loss.")
distill_alpha_end: Optional[float] = Field(None, description="Target alpha at end of training. None keeps alpha fixed.")
distill_alpha_schedule: Literal["constant", "linear", "cosine"] = Field(
"constant", description="Schedule type for alpha annealing ('constant', 'linear', or 'cosine')."
"constant",
description="Schedule type for alpha annealing ('constant', 'linear', or 'cosine').",
)
distill_temperature_end: Optional[float] = Field(
None, description="Target temperature at end of training. None keeps temperature fixed."
None,
description="Target temperature at end of training. None keeps temperature fixed.",
)
distill_temperature_schedule: Literal["constant", "linear", "cosine"] = Field(
"constant", description="Schedule type for temperature annealing ('constant', 'linear', or 'cosine')."
"constant",
description="Schedule type for temperature annealing ('constant', 'linear', or 'cosine').",
)
distill_beta_end: Optional[float] = Field(
None, description="Target beta_feature at end of training. None keeps beta fixed."
None,
description="Target beta_feature at end of training. None keeps beta fixed.",
)
distill_beta_schedule: Literal["constant", "linear", "cosine"] = Field(
"constant", description="Schedule type for beta annealing ('constant', 'linear', or 'cosine')."
"constant",
description="Schedule type for beta annealing ('constant', 'linear', or 'cosine').",
)

# --- Learn to init related parameters --
Expand All @@ -1463,11 +1490,13 @@ class Distillation(BaseModel):
)

attn_module_name: Optional[str] = Field(
None, description="Attention nnx module attribute name to augment with LTI logic"
None,
description="Attention nnx module attribute name to augment with LTI logic",
)

lti_layer_indices: Optional[list[int]] = Field(
None, description="List of layer indices to apply LTI modifications. If None, applied to all layers."
None,
description="List of layer indices to apply LTI modifications. If None, applied to all layers.",
)
# ---------------------------------------

Expand Down Expand Up @@ -1532,11 +1561,13 @@ class DilocoParams(BaseModel):
diloco_outer_lr: float = Field(0.3, description="learning rate for outer optimizer.")
diloco_outer_momentum: float = Field(0.9, description="momentum for outer optimizer.")
dcn_bandwidth_limit: str = Field(
"", description="Programmatic DCN egress bandwidth limit (e.g., '28gbit'). Empty means no limit."
"",
description="Programmatic DCN egress bandwidth limit (e.g., '28gbit'). Empty means no limit.",
)
dcn_bandwidth_burst: str = Field("10mb", description="Burst size for Token Bucket Filter (TBF) traffic shaping.")
dcn_bandwidth_latency: str = Field(
"50ms", description="Latency threshold for Token Bucket Filter (TBF) traffic shaping."
"50ms",
description="Latency threshold for Token Bucket Filter (TBF) traffic shaping.",
)
dcn_bandwidth_interface: str = Field("eth0", description="Network interface to apply bandwidth limits on.")

Expand Down Expand Up @@ -1829,7 +1860,8 @@ class Profiling(BaseModel):
tpu_num_chips_to_profile_per_task: int = Field(1, description="Specifies the number of TPU chips to profile per task.")
tpu_num_sparse_cores_to_trace: int = Field(2, description="Specifies the number of TPU chips to profile per task.")
tpu_num_sparse_core_tiles_to_trace: int = Field(
1, description="Specifies the number of tiles within each sparse core to trace on the TPU."
1,
description="Specifies the number of tiles within each sparse core to trace on the TPU.",
)
xprof_tpu_power_trace_level: XProfTPUPowerTraceMode = Field(
XProfTPUPowerTraceMode.POWER_TRACE_NONE,
Expand Down Expand Up @@ -2770,7 +2802,11 @@ def validate_and_set_hlo_dump_defaults():
)
for param_name, schedule, end_value in [
("distill_alpha", self.distill_alpha_schedule, self.distill_alpha_end),
("distill_temperature", self.distill_temperature_schedule, self.distill_temperature_end),
(
"distill_temperature",
self.distill_temperature_schedule,
self.distill_temperature_end,
),
("distill_beta", self.distill_beta_schedule, self.distill_beta_end),
]:
if schedule != "constant" and end_value is None:
Expand Down Expand Up @@ -2827,7 +2863,10 @@ def get_num_target_devices():

# Check for AQT deprecation warning
if self.quantization and not self.use_qwix_quantization:
if self.quantization not in ("fp8", "nanoo_fp8") and not self.quantization.startswith("te_"):
if self.quantization not in (
"fp8",
"nanoo_fp8",
) and not self.quantization.startswith("te_"):
logger.warning(
"WARNING: AQT quantization is deprecated and will be removed in a future release. "
"Please migrate to Qwix by setting use_qwix_quantization=True."
Expand Down Expand Up @@ -2935,6 +2974,7 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de
"query_proj",
"key_proj",
"value_proj",
"kv_proj",
"query_wa_proj",
"kv_wa_proj",
"mla_kv",
Expand Down
Loading
Loading