| Name | Modified | Size | Downloads / Week |
|---|---|---|---|
| Parent folder | |||
| maxtext-v0.2.5 source code.tar.gz | 2026-10-01 | 43.1 MB | |
| maxtext-v0.2.5 source code.zip | 2026-10-01 | 44.5 MB | |
| README.md | 2026-10-01 | 16.9 kB | |
| Totals: 3 Items | 87.6 MB | 0 | |
Changes
Models
- DeepSeek-V3 Lineage Integration: Integrated Lineage DeepSeek-V3 (
dsv3.py) into MaxText, storing routed expert output weights in Lineage's layout and sowing router bias updates along theuse_lineagepath. - DeepSeek-V4 & Multi-Token Prediction: Supported DeepSeek-V4 checkpoint conversion and parameter mapping with 2D transposed shape alignment, added DeepSeek-V4 indexer loss, and implemented separate final normalization for Multi-Token Prediction (MTP).
- Weaver Architecture: Added Weaver Mixture-of-Transformers block and standardized Weaver model naming across configs.
- Qwen3.5 & Qwen3-Next: Added
qwen3.5-35b-a3b-fp8andqwen3.5-397b-a17b-fp8model configs with HuggingFace parameter mapping and weight conversion, and optimized Qwen3-Next shared expert construction to only instantiate whenshared_experts > 0. - Cosmos 3 Bring-Up: Added model definitions and bring-up support for Cosmos 3 Nano and Cosmos 3 Super Reasoner.
- Normalization Enhancements: Added RMSNorm normalization layer support for Apple Envy decoder blocks and OLMo 3 MLP blocks (
DecoderBlockType.OLMO3). - GPT-OSS Attention Sink: Added attention-sink support for GPT-OSS models.
- M3 Core Architecture: Added m3 package scaffolding with architecture rules and layout verification tests, introducing m3 core RoPE implementation, mesh creation, and sharding rules.
Pre-Training
- JAX Flash Attention Optimization: Enhanced JAX Flash Attention kernel with a specialized high-performance execution path for long sequences (splitting at a 4,096 sequence length threshold).
- DeepSeek-V4 CSA & Splash Attention: Added fused Pallas TPU StreamIndex score kernel, QK attention head chunking for Compressed Self-Attention (CSA) memory reduction, and HCA static compilation for Splash Attention on DeepSeek-V4.
- Context Parallelism (CP): Sharded GatedDeltaNet sequences under
ici_context_parallelism, enabled native block-sparse execution for indexer masks under CP, removed all-to-all communication overhead from TSP/CP segmentation masks, and split embedding logical names in attention modules. - Multimodal RoPE (MRoPE): Eliminated redundant JAX recompilations in
get_mrope_input_positionsusing inline NumPy operations, and normalized MRoPE position shapes for vLLM across hybrid cache utilities and embeddings. - DiLoCo Distributed Training: Added comprehensive documentation and tutorials for DiLoCo pre-training (core concepts, diloco, tutorials), and migrated DiLoCo launch scripts and tutorials from XPK to Cluster Toolkit (diloco pretraining).
- Quantized Multi-Token Prediction: Enabled quantized Multi-Token Prediction (MTP) for accelerated pre-training.
Post-Training
- Reinforcement Learning (GRPO) Infrastructure: Plumbed GRPO flags to
AgenticGRPOLearner, added RL sequence packing and vLLM block-size configuration knobs, addedfree_kv_cache_during_weight_sync,reward_num_workers, andgc_collect_after_weight_syncsettings, madeextract_answerlinear in response length, and added a Gemma 4 26B GRPO tutorial (post training index, rl gemma4 26b, rl gemma4 e4b). - LoRA & Distillation Workflows: Refactored LoRA resharding to modernize NNX variable access and graph traversal, migrated full fine-tuning scripts from TFDS to Grain (full finetuning), added launch scripts and configs for Distillation Training, and integrated MaxText's trainer backend into a distributed GSM8K example.
- vLLM Rollout Enhancements: Added Mamba prefix caching for Qwen3 GDN layers in vLLM, indexed KV caches via
layer_name_to_kvcache_indexin the vLLM adapter, and implemented MoE 128-lane chunking and padding for GMM layout.
Multimodal
- Multimodal SFT Pipeline: Added Grain-based multimodal Supervised Fine-Tuning (SFT) data processing pipeline (multimodal), end-to-end video SFT supporting variable frame sizes and durations, and ChartNet SFT configuration with string response handling in multimodal input pipelines.
- Omni Pipeline & Qwen3-VL Inference: Added custom multimodal processor and end-to-end pipeline for Omni multimodal models, and enabled
vllm_decodeinference support for Qwen3-VL (inference).
Performance
- Explicit Sharding & ZeRO-1 Expansion: Onboarded explicit sharding and ZeRO-1 gradient accumulation across Qwen3.5, Qwen3-Next, Qwen3, Qwen2, Kimi-K2, Mistral, Mixtral, and Gemma model families; supported partitioned optimizer state sharding propagation and MaskedNode filtering in ZeRO-1; deferred data-parallel gradient all-reduce to
update()under GA with backward cotangents scaled to GA=1. - Sharded Muon Optimizer: Integrated a sharding-aware, distributed Muon optimizer into MaxText with generalized weight dimension extraction, support for Qwen3 and GPT-OSS MoE blocks, and a configuration flag to govern applying Muon to MoE routers.
- FP8 Quantization & Dequantization: Implemented an FP8 weight-only dynamic dequantization engine with direct FP8 and scale tensor ingestion in
to_maxtext, native FP8 inference (serve_fp8_weight) for DenseGeneral and GMM v2, weight-gradient-arm cotangent calibration (drhs_grad_quantization_calibration_method), quantized logits projection (quantization), pre-quantized weights for TPU inference fused MoE kernels, and options to keep MoE router projections unquantized for stability (moe configuration, quantization). - Quantized Ring-of-Experts Pipeline: Enabled token-quantized all-gather pipelines with Ring of Experts and added quantized all-gather for the Ring of Experts combine backward (moe configuration, quantization).
- MoE Dropless Fallback & Buffer Resilience: Implemented configurable MoE dropless fallback (
moe_dropless_fallback: None|step|layer) with memory-neutral step fallback, first-phase ragged buffers, host-driven step replay retry mechanisms for buffer overflows, evaluation-specific buffer scaling (eval_ragged_buffer_factor), and evaluation tile sizes and QKV layouts (moe configuration). - MoE Routing & Kernel Optimizations: Added dense custom VJPs for router top-k selection (eliminating backward scatters), optimized top-2 expert group score computation, aligned sigmoid-router auxiliary losses and expert bias updates with Megatron-LM, added collective psum reductions for expert bias load balancing, and supported extracting and replaying router expert decisions between inference and training (moe configuration).
- MoE Kernels & Sharding: Integrated TransformerEngine (TE) MoEBlock, supported Tokamax GMM v2 heuristic tiling functions, forwarded TPU inference MoE kernel knobs, saved routed expert inputs across remat boundaries via
moe_x_sorted, added nested scanning over hybrid attention with per-layer remat in Qwen3-Next, and added explicit tracking forshard_embed_moe_on_fsdp(moe configuration). - SparseCore & Pallas Accelerations: Added
moe_pin_sparse_core_all_gathersto offload FSDP and Expert Parallel all-gathers to SparseCores, loweredplsc.bitcasttotpu.bitcastwhen layout passes are enabled, and refactored the MaxText mHC kernel to utilizeemit_pipeline.
Checkpointing / Goodput
- Orbax v1 Checkpointing: Refactored MaxText's checkpointing architecture to natively utilize the Orbax v1 API (retiring legacy v0 implementations), added dequantize-on-load parameter restoration, and added support for restoring checkpoint weights stored outside the standard
"params"collection. - Goodput & Pathways Persistence: Added
goodput_job_nameconfiguration option for Goodput monitoring, registered Pathways persistence array handlers with bounded host staging, and introduced standalone checkpointer benchmarking loop features and configuration flags. - Checkpointing Documentation: Migrated multi-tier and emergency checkpointing guides from XPK to Cluster Toolkit (emergency checkpointing, multi tier checkpointing).
Usability
- MLPerf Training Compliance & Optimizations: Added full MLPerf training logging compliance, accelerated step-0 initialization time via per-file
c4_mlperftrain sharding and data-free setup, added host-local eval batch caching, implemented continuous stream chunking for C4 MLPerf datasets, and added thedsv3-mlperf-4kcustom mesh and sharding rule. - Data Pipeline Enhancements: Added Megatron MMap data format support for Grain data loaders with an integrated mmap index builder (data input pipeline, data input grain, data input megatron mmap), allowed Grain to ingest TFDS configurations, optimized data input pipelines for evaluation, and set default
dataset_typetosynthetic. - Documentation & Cluster Toolkit Migration: Migrated core navigation, getting started, and execution guides from XPK to Cluster Toolkit (getting started, install maxtext, architecture overview), updated container image registry references from GCR to Artifact Registry (run maxtext elastic training, run maxtext via pathways, lora on multi host), and published TPU pre-training and post-training Docker images for MaxText in the build tutorial (build maxtext).
- Dependencies & Developer Tooling: Upgraded JAX dependencies to 0.11.1 (update dependencies), Tokamax to 0.0.14, and DrJAX to 0.2.1; added custom GCS path support for ML Diagnostics telemetry; added Gemini CLI development skills; streamlined multi-worker XLA dumping (
--dump_xla_all_hosts); and defined default rule constants intypes.pywith base configuration parity testing.
Bug Fixes
- Fixed FP8 and quantized MoE numerical paths across dense and sparse matmul implementations.
- Fixed MoE weight tracking, vLLM HBM memory allocation, and restored
woin Tunix's zero-pad key set during RL rollouts. - Fixed RLConfig mesh initialization AttributeError and resolved defects on the Tunix adapter and engine packing path.
- Fixed vLLM sampler initialization failures and corrected
cache_headsandpaged_kv_headslogical axis order invllm.yml. - Unblocked Qwen3.5 hybrid decode through the vLLM adapter and fixed the standalone torchax converter for TPU inference rollouts.
- Fixed standalone checkpointer restore loop to use Orbax v1 and ensured fatal checkpointing errors raise
RuntimeError. - Resolved Pyrefly typing errors related to JAX scalar types across trainers and Megablox kernels.
- Derived
cp_sizedynamically from active logical axis rules and computed local block padding for dynamic Splash indexer masks under Context Parallelism. - Fixed GDN Mamba block table indexing under prefix caching and resolved argument mismatches in compact Mamba block override calls.
- Fixed dropped routing-weight gradients in
ring_ragged_unsortand aligned dropless-replay remat overrides with the NNX decoder. - Configured layout passes on SparseCore Pallas ragged gather and ragged gather-reduce kernels.
- Eliminated XLA backward adjoint interior padding in
reorder_sequence. - Fixed duplicate XLA compilation of
jit(train_step). - Fixed vocab tiling hidden-state cotangent scaling.
- Fixed FP8 custom-gradient state overwrite during NNX training.
- Fixed NNX throughput regression when running with
scan_layers=False. - Fixed XLA fusion breakage and numerical drift in Qwen3 under AUTO sharding mode.
- Fixed edge cases in
AttentionOp.generate_attention_maskfor DeepSeek-V4 compressed attention across chunked prefill, packed sequences, and standalone environments. - Fixed Qwen3-VL tokenizer type mismatch and resolved Grain
IterDatasetcrashes on multimodal data. - Ensured
weight_dtypeis propagated correctly in Qwen3-Next sparse MoE shared expert gate. - Fixed EOS token detection and model routing for Qwen vision models in multimodal evaluation.
- Fixed GCS path resolution by skipping redundant
exists()checks on Google Cloud Storage paths. - Allowed warm starts from
load_parameters_pathwithout requiring checkpointing enabled. - Fixed MoE Data Parallel reduce-scatter and logical axis sharding in the vLLM serving path.
- Fixed Ahead-Of-Time (AOT) compilation for ragged kernels and supported internal compilation and orchestration workflows.
- Fixed C4 MLPerf evaluation input pipeline handling for fractional evaluation batch sizes.
- Resolved Grain dataloader starvation and process leak on slice recovery.
- Fixed chat template formatting and SFT prompt/completion masking for Gemma 4 reasoning.
- Fixed
ImportErroron optionaltensorflow_textin post-training workflows. - Explicitly set
enable_checkpointingwhen loading state in DiLoCo tests. - Fixed
AttributeErrorin MHC-lite during shape evaluation and sharding construction. - Fixed inference query head rule mapping and aligned QKV batch dimensions for cuDNN flash attention.
Deprecations
- Flax Linen Removal: Removed legacy Flax Linen module code, dead code, and shared-primitive imports, deprecating and removing
pure_nnx,enable_nnx, andpure_nnx_decoderconfiguration flags, Linen decoder/attention layers, and*_as_linenmodel wrappers across trainers, inference dispatch, quantization, and sharding (distillation, knowledge distillation, lora). - Replaced deprecated
google-cloud-sdkpackage withgoogle-cloud-cliin Dockerfiles (run maxtext via xpk, rl on multi host, sft on multi host). - Made TensorFlow an optional dependency in requirements (update dependencies, data input tfds, install maxtext).