Download Latest Version maxtext-v0.2.5 source code.zip (44.5 MB) Google Add to Preferred Sources
Home / maxtext-v0.2.5
Name Modified Size InfoDownloads / 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 the use_lineage path.
  • 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-fp8 and qwen3.5-397b-a17b-fp8 model configs with HuggingFace parameter mapping and weight conversion, and optimized Qwen3-Next shared expert construction to only instantiate when shared_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_positions using 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, added free_kv_cache_during_weight_sync, reward_num_workers, and gc_collect_after_weight_sync settings, made extract_answer linear 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_index in 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_decode inference 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 for shard_embed_moe_on_fsdp (moe configuration).
  • SparseCore & Pallas Accelerations: Added moe_pin_sparse_core_all_gathers to offload FSDP and Expert Parallel all-gathers to SparseCores, lowered plsc.bitcast to tpu.bitcast when layout passes are enabled, and refactored the MaxText mHC kernel to utilize emit_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_name configuration 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_mlperf train sharding and data-free setup, added host-local eval batch caching, implemented continuous stream chunking for C4 MLPerf datasets, and added the dsv3-mlperf-4k custom 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_type to synthetic.
  • 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 in types.py with 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 wo in 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_heads and paged_kv_heads logical axis order in vllm.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_size dynamically 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_unsort and 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_mask for DeepSeek-V4 compressed attention across chunked prefill, packed sequences, and standalone environments.
  • Fixed Qwen3-VL tokenizer type mismatch and resolved Grain IterDataset crashes on multimodal data.
  • Ensured weight_dtype is 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_path without 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 ImportError on optional tensorflow_text in post-training workflows.
  • Explicitly set enable_checkpointing when loading state in DiLoCo tests.
  • Fixed AttributeError in MHC-lite during shape evaluation and sharding construction.
  • Fixed inference query head rule mapping and aligned QKV batch dimensions for cuDNN flash attention.

Deprecations

Source: README.md, updated 2026-10-01