diff --git a/.dockerignore b/.dockerignore index e567ea2ff6..fe3d4214b5 100644 --- a/.dockerignore +++ b/.dockerignore @@ -1,2 +1,4 @@ .git maxtext_venv +.venv +venv13 diff --git a/run_qwen3_80b_aot.sh b/run_qwen3_80b_aot.sh new file mode 100755 index 0000000000..c69c617ca2 --- /dev/null +++ b/run_qwen3_80b_aot.sh @@ -0,0 +1,156 @@ +#!/bin/bash +set -e + +# Activate Python virtual environment (~/.venv) +source /usr/local/google/home/chengnuojin/.venv/bin/activate + +# Execute in a subshell to ensure environment recovery automatically +( + # --- 1. Set XLA Flags --- + XLA_FLAGS_ARRAY=( + "--xla_msa_enable_sync_slice_replacement=false" + "--xla_tpu_enable_sparse_core_collective_offload_2d_all_gather=true" + "--xla_tpu_enable_sparse_core_collective_offload_reduce_scatter=true" + "--xla_msa_enable_sync_copy_replacement=false" + "--xla_tpu_scoped_vmem_limit_kib=81000" + "--xla_tpu_enable_sparse_core_collective_offload_all_gather=true" + "--xla_tpu_enable_sparse_core_collective_offload_all_reduce=true" + "--xla_tpu_enable_concurrent_sparse_core_offloading=true" + "--xla_tpu_enable_sparse_core_offload_queuing_in_lhs=true" + "--xla_tpu_enable_layer_scheduler_for_dependent_collectives=true" + "--xla_tpu_use_single_sparse_core_for_all_gather_offload=true" + "--xla_tpu_sparse_core_all_gather_latency_multiplier=1" + "--xla_tpu_sparse_core_reduce_scatter_latency_multiplier=3" + "--xla_tpu_offload_gather_to_sparsecore=true" + "--xla_tpu_dvfs_p_state=7" + "--xla_tpu_disable_sparse_core_collective_offload_remover=true" + "--xla_tpu_use_tc_device_shape_on_sc=true" + "--xla_sc_enable_instruction_fusion=false" + "--xla_sc_disable_megacore_partitioning=true" + "--xla_tpu_enable_async_collective_fusion=true" + "--xla_tpu_overlap_compute_collective_tc=true" + "--xla_tpu_enable_async_collective_fusion_multiple_steps=true" + "--xla_tpu_enable_async_collective_fusion_fuse_all_gather=false" + "--xla_tpu_enable_async_collective_fusion_fuse_reduce_scatter=false" + "--xla_tpu_enable_async_collective_fusion_fuse_all_reduce=false" + "--xla_tpu_enable_latency_hiding_scheduler=true" + "--xla_latency_hiding_scheduler_rerun=10" + "--xla_tpu_all_gather_collective_matmul_mode=post_spmd_conservative" + "--xla_tpu_reduce_scatter_collective_matmul_mode=post_spmd_conservative" + "--xla_latency_hiding_scheduler_enable_selective_resources=true" + "--xla_tpu_enable_ilp_latency_hiding_scheduler=true" + "--xla_tpu_enable_all_experimental_scheduler_features=true" + "--xla_tpu_enable_scheduler_memory_pressure_tracking=true" + "--xla_tpu_host_transfer_overlap_limit=4" + "--xla_tpu_aggressive_opt_barrier_removal=ENABLED" + "--xla_lhs_prioritize_async_depth_over_stall=ENABLED" + "--xla_tpu_enable_ag_backward_pipelining=true" + "--xla_should_allow_loop_variant_parameter_in_chain=ENABLED" + "--xla_should_add_loop_invariant_op_in_chain=ENABLED" + "--xla_max_concurrent_host_send_recv=100" + "--xla_tpu_scheduler_percent_shared_memory_limit=150" +) + export LIBTPU_INIT_ARGS="${XLA_FLAGS_ARRAY[*]}" + + # --- 2. Export Required Environment Variables --- + export PYTHONPATH=$PWD/src:$PWD/src/maxtext/src:$PYTHONPATH + export JAX_PLATFORMS='cpu' + export ENABLE_PJRT_COMPATIBILITY='true' + + # --- 3. Configuration --- + TIMESTAMP=$(date +%m%d%H%M%S) + export MODEL_NAME="qwen3-next-80b-a3b" + export BASE_OUTPUT_DIR="gs://chengnuojin-maxtext-logs/qwen3-next-80b-profiles/run-${TIMESTAMP}" + + # --- 4. Run Train Compile (AOT Compilation) --- + MAXTEXT_ARGS_ARRAY=( + "compile_xla_flags=${LIBTPU_INIT_ARGS}" + "model_name=${MODEL_NAME}" + "base_output_directory=${BASE_OUTPUT_DIR}" + "run_name=param-3" + "log_config=false" + "debug_sharding=false" + "ragged_gather_reduce_fallback=false" + "compile_topology=v6e-256" + "compile_topology_num_slices=1" + "dataset_type=synthetic" + "dataset_name=synthetic" + "dtype=bfloat16" + "allow_split_physical_axes=True" + "ici_expert_parallelism=4" + "use_ring_of_experts=True" + "custom_mesh=hybrid_ring_64x4" + "use_ragged_sort=True" + "use_random_routing=True" + "num_moe_token_chunks=2" + "per_device_batch_size=8" + "opt_type=adamw" + "max_target_length=2048" + "ragged_buffer_factor=1.5" + "remat_policy=custom" + "reuse_example_batch=1" + "decoder_layer_input=device" + "ici_fsdp_parallelism=-1" + "steps=15" + "sa_q_layout=SEQ_MINOR" + "sa_k_layout=HEAD_DIM_MINOR" + "sa_v_layout=HEAD_DIM_MINOR" + "sa_block_q=1024" + "sa_block_kv=1024" + "sa_block_kv_compute=512" + "sa_block_q_dkv=1024" + "sa_block_kv_dkv=1024" + "sa_block_kv_dkv_compute=1024" + "sa_fuse_reciprocal=false" + "use_splash_scheduler=true" + "sa_use_base2_exp=true" + "dq_reduction_steps=3" + "hardware=tpu" + "skip_jax_distributed_system=True" + "attention=flash" + "use_tokamax_splash=True" + "sa_use_fused_bwd_kernel=True" + "sparse_matmul=True" + "megablox=True" + "wi_tile_fwd_batch_seq=128" + "wi_tile_dlhs_batch_seq=128" + "wi_tile_drhs_batch_seq=128" + "wo_tile_fwd_batch_seq=128" + "wo_tile_dlhs_batch_seq=128" + "wo_tile_drhs_batch_seq=128" + "wi_tile_fwd_embed_dim=3072" + "wi_tile_fwd_mlp_dim=1536" + "wi_tile_dlhs_embed_dim=3072" + "wi_tile_dlhs_mlp_dim=1536" + "wi_tile_drhs_embed_dim=3072" + "wi_tile_drhs_mlp_dim=1536" + "wo_tile_fwd_embed_dim=3072" + "wo_tile_fwd_mlp_dim=1536" + "wo_tile_dlhs_embed_dim=3072" + "wo_tile_dlhs_mlp_dim=1536" + "wo_tile_drhs_embed_dim=3072" + "wo_tile_drhs_mlp_dim=1536" + "use_tokamax_gmm=True" + "use_gmm_v2=True" + "optimizer_memory_host_offload=False" + "parameter_memory_host_offload=False" + "enable_checkpointing=False" + "async_checkpointing=False" + "tokenizer_type=tiktoken" + "tokenizer_path=tokenizer_74B/" + "use_gdn_kernel=True" + "use_hybrid_gdn=True" + "profiler=xplane" + "profiler_steps=5" + "skip_first_n_steps_for_profiler=2" + "enable_tpu_profiling_options=True" + "upload_all_profiler_results=False" + ) + + echo "========================================================================" + echo "Running MaxText AOT train_compile for ${MODEL_NAME} on v6e-256" + echo "========================================================================" + + rm -f /tmp/libtpu_lockfile 2>/dev/null || true + python3 -m maxtext.trainers.pre_train.train_compile src/maxtext/configs/base.yml "${MAXTEXT_ARGS_ARRAY[@]}" +) diff --git a/run_qwen3_80b_aot_fsdp.sh b/run_qwen3_80b_aot_fsdp.sh new file mode 100755 index 0000000000..9a22c80fc1 --- /dev/null +++ b/run_qwen3_80b_aot_fsdp.sh @@ -0,0 +1,176 @@ +#!/bin/bash +set -e + +# Activate Python virtual environment (~/.venv) +source /usr/local/google/home/chengnuojin/.venv/bin/activate + +# Execute in a subshell to ensure environment recovery automatically +( + # --- 1. Set XLA Flags --- + # XLA_FLAGS_ARRAY=( + # "--xla_msa_enable_sync_slice_replacement=false" + # "--xla_tpu_enable_sparse_core_collective_offload_2d_all_gather=true" + # "--xla_tpu_enable_sparse_core_collective_offload_reduce_scatter=true" + # "--xla_msa_enable_sync_copy_replacement=false" + # "--xla_tpu_scoped_vmem_limit_kib=81000" + # "--xla_tpu_enable_sparse_core_collective_offload_all_gather=true" + # "--xla_tpu_enable_sparse_core_collective_offload_all_reduce=true" + # "--xla_tpu_offload_gather_to_sparsecore=true" + # "--xla_tpu_dvfs_p_state=7" + # "--xla_tpu_disable_sparse_core_collective_offload_remover=true" + # "--xla_tpu_enable_async_collective_fusion=true" + # "--xla_tpu_overlap_compute_collective_tc=true" + # "--xla_tpu_enable_async_collective_fusion_multiple_steps=true" + # "--xla_tpu_enable_latency_hiding_scheduler=true" + # "--xla_latency_hiding_scheduler_rerun=10" + # "--xla_tpu_all_gather_collective_matmul_mode=post_spmd_conservative" + # "--xla_tpu_reduce_scatter_collective_matmul_mode=post_spmd_conservative" + # "--xla_latency_hiding_scheduler_enable_selective_resources=true" + # "--xla_tpu_enable_ilp_latency_hiding_scheduler=true" + # "--xla_tpu_enable_all_experimental_scheduler_features=true" + # "--xla_tpu_enable_scheduler_memory_pressure_tracking=true" + # "--xla_tpu_host_transfer_overlap_limit=4" + # "--xla_tpu_aggressive_opt_barrier_removal=ENABLED" + # "--xla_lhs_prioritize_async_depth_over_stall=DISABLED" + # "--xla_tpu_enable_ag_backward_pipelining=true" + # "--xla_should_allow_loop_variant_parameter_in_chain=ENABLED" + # "--xla_should_add_loop_invariant_op_in_chain=ENABLED" + # "--xla_max_concurrent_host_send_recv=100" + # "--xla_tpu_scheduler_percent_shared_memory_limit=50" + # ) + XLA_FLAGS_ARRAY=( + "--xla_msa_enable_sync_slice_replacement=false" + "--xla_tpu_enable_sparse_core_collective_offload_2d_all_gather=true" + "--xla_tpu_enable_sparse_core_collective_offload_reduce_scatter=true" + "--xla_msa_enable_sync_copy_replacement=false" + "--xla_tpu_scoped_vmem_limit_kib=81000" + "--xla_tpu_enable_sparse_core_collective_offload_all_gather=true" + "--xla_tpu_enable_sparse_core_collective_offload_all_reduce=true" + "--xla_tpu_enable_concurrent_sparse_core_offloading=true" + "--xla_tpu_enable_sparse_core_offload_queuing_in_lhs=true" + "--xla_tpu_enable_layer_scheduler_for_dependent_collectives=true" + "--xla_tpu_use_single_sparse_core_for_all_gather_offload=true" + "--xla_tpu_sparse_core_all_gather_latency_multiplier=1" + "--xla_tpu_sparse_core_reduce_scatter_latency_multiplier=3" + "--xla_tpu_offload_gather_to_sparsecore=true" + "--xla_tpu_dvfs_p_state=7" + "--xla_tpu_disable_sparse_core_collective_offload_remover=true" + "--xla_tpu_use_tc_device_shape_on_sc=true" + "--xla_sc_enable_instruction_fusion=false" + "--xla_sc_disable_megacore_partitioning=true" + "--xla_tpu_enable_async_collective_fusion=true" + "--xla_tpu_overlap_compute_collective_tc=true" + "--xla_tpu_enable_async_collective_fusion_multiple_steps=true" + "--xla_tpu_enable_async_collective_fusion_fuse_all_gather=false" + "--xla_tpu_enable_async_collective_fusion_fuse_reduce_scatter=false" + "--xla_tpu_enable_async_collective_fusion_fuse_all_reduce=false" + "--xla_tpu_enable_latency_hiding_scheduler=true" + "--xla_latency_hiding_scheduler_rerun=10" + "--xla_tpu_all_gather_collective_matmul_mode=post_spmd_conservative" + "--xla_tpu_reduce_scatter_collective_matmul_mode=post_spmd_conservative" + "--xla_latency_hiding_scheduler_enable_selective_resources=true" + "--xla_tpu_enable_ilp_latency_hiding_scheduler=true" + "--xla_tpu_enable_all_experimental_scheduler_features=true" + "--xla_tpu_enable_scheduler_memory_pressure_tracking=true" + "--xla_tpu_host_transfer_overlap_limit=4" + "--xla_tpu_aggressive_opt_barrier_removal=ENABLED" + "--xla_lhs_prioritize_async_depth_over_stall=ENABLED" + "--xla_tpu_enable_ag_backward_pipelining=true" + "--xla_should_allow_loop_variant_parameter_in_chain=ENABLED" + "--xla_should_add_loop_invariant_op_in_chain=ENABLED" + "--xla_max_concurrent_host_send_recv=100" + "--xla_tpu_scheduler_percent_shared_memory_limit=150" +) + export LIBTPU_INIT_ARGS="${XLA_FLAGS_ARRAY[*]}" + + # --- 2. Export Required Environment Variables --- + export PYTHONPATH=$PWD/src:$PWD/src/maxtext/src:$PYTHONPATH + export JAX_PLATFORMS='cpu' + export ENABLE_PJRT_COMPATIBILITY='true' + + # --- 3. Configuration --- + TIMESTAMP=$(date +%m%d%H%M%S) + export MODEL_NAME="qwen3-next-80b-a3b" + export BASE_OUTPUT_DIR="gs://chengnuojin-maxtext-logs/qwen3-next-80b-profiles/run-${TIMESTAMP}" + + # --- 4. Run Train Compile (AOT Compilation) --- + MAXTEXT_ARGS_ARRAY=( + "compile_xla_flags=${LIBTPU_INIT_ARGS}" + "model_name=${MODEL_NAME}" + "base_output_directory=${BASE_OUTPUT_DIR}" + "run_name=param-3" + "log_config=false" + "debug_sharding=false" + "compile_topology=v6e-256" + "compile_topology_num_slices=1" + "dataset_type=synthetic" + "dataset_name=synthetic" + "dtype=bfloat16" + "allow_split_physical_axes=False" + "ici_expert_parallelism=1" + "per_device_batch_size=6" + "opt_type=adamw" + "max_target_length=2048" + "remat_policy=custom" + "reuse_example_batch=1" + "decoder_layer_input=device" + "ici_fsdp_parallelism=-1" + "steps=20" + "sa_q_layout=SEQ_MINOR" + "sa_k_layout=HEAD_DIM_MINOR" + "sa_v_layout=HEAD_DIM_MINOR" + "sa_block_q=2048" + "sa_block_kv=2048" + "sa_block_kv_compute=1024" + "sa_block_q_dkv=2048" + "sa_block_kv_dkv=2048" + "sa_block_kv_dkv_compute=1024" + "hardware=tpu" + "skip_jax_distributed_system=True" + "attention=flash" + "use_tokamax_splash=True" + "sa_use_fused_bwd_kernel=True" + "sparse_matmul=True" + "megablox=True" + "wi_tile_fwd_batch_seq=64" + "wi_tile_dlhs_batch_seq=64" + "wi_tile_drhs_batch_seq=64" + "wo_tile_fwd_batch_seq=64" + "wo_tile_dlhs_batch_seq=64" + "wo_tile_drhs_batch_seq=64" + "wi_tile_fwd_embed_dim=3072" + "wi_tile_fwd_mlp_dim=1536" + "wi_tile_dlhs_embed_dim=3072" + "wi_tile_dlhs_mlp_dim=1536" + "wi_tile_drhs_embed_dim=3072" + "wi_tile_drhs_mlp_dim=1536" + "wo_tile_fwd_embed_dim=3072" + "wo_tile_fwd_mlp_dim=1536" + "wo_tile_dlhs_embed_dim=3072" + "wo_tile_dlhs_mlp_dim=1536" + "wo_tile_drhs_embed_dim=3072" + "wo_tile_drhs_mlp_dim=1536" + "use_tokamax_gmm=True" + "use_gmm_v2=True" + "optimizer_memory_host_offload=False" + "parameter_memory_host_offload=False" + "enable_checkpointing=False" + "async_checkpointing=False" + "tokenizer_type=tiktoken" + "tokenizer_path=tokenizer_74B/" + "use_gdn_kernel=True" + "use_hybrid_gdn=True" + "profiler=xplane" + "profiler_steps=2" + "skip_first_n_steps_for_profiler=2" + "enable_tpu_profiling_options=True" + "upload_all_profiler_results=False" + ) + + echo "========================================================================" + echo "Running MaxText AOT train_compile for ${MODEL_NAME} on v6e-256" + echo "========================================================================" + + rm -f /tmp/libtpu_lockfile 2>/dev/null || true + python3 -m maxtext.trainers.pre_train.train_compile src/maxtext/configs/base.yml "${MAXTEXT_ARGS_ARRAY[@]}" +) diff --git a/run_qwen3_80b_xpk.sh b/run_qwen3_80b_xpk.sh new file mode 100755 index 0000000000..0466b60734 --- /dev/null +++ b/run_qwen3_80b_xpk.sh @@ -0,0 +1,204 @@ +#!/bin/bash +set -e + +# Activate Python virtual environment +source /usr/local/google/home/chengnuojin/.venv/bin/activate + +# --- Environment Variables --- +export PROJECT_ID="tpu-prod-env-one-vm" +export CLUSTER_NAME="bodaborg-v6e-256-lcscld-c" +export ZONE="southamerica-west1-a" + +# --- Configuration & Automated Image Build --- +TIMESTAMP=$(date +%m%d%H%M%S) +export WORKLOAD_IMAGE="gcr.io/tpu-prod-env-one-vm/param3_21jul:chengnuojin_${TIMESTAMP}" +export WORKLOAD_NAME="chengnuojin-qn80b-${TIMESTAMP}" +export DEVICE_TYPE="v6e-256" +export NUM_SLICES=1 +export PRIORITY="very-high" +export NUM_STEPS=15 +export MODEL_NAME="qwen3-next-80b-a3b" +export BASE_OUTPUT_DIR="gs://chengnuojin-maxtext-logs/qwen3-next-80b-profiles/run-${TIMESTAMP}" + +echo "========================================================================" +echo "Building and uploading Docker runner image from /usr/local/google/home/chengnuojin/maxtext" +echo "Target Image: ${WORKLOAD_IMAGE}" +echo "========================================================================" + +( + cd /usr/local/google/home/chengnuojin/maxtext && \ + if ! docker image inspect maxtext_base_image &> /dev/null; then + echo "Local image 'maxtext_base_image' not found. Pulling gcr.io/tpu-prod-env-one-vm/param3_21jul:latest and tagging as maxtext_base_image..." + docker pull gcr.io/tpu-prod-env-one-vm/param3_21jul:latest + docker tag gcr.io/tpu-prod-env-one-vm/param3_21jul:latest maxtext_base_image + fi && \ + CLOUD_IMAGE_NAME="${WORKLOAD_IMAGE}" \ + BASE_IMAGE="gcr.io/tpu-prod-env-one-vm/param3_21jul:latest" \ + bash src/dependencies/scripts/docker_upload_runner.sh +) + +echo "Docker image upload complete: ${WORKLOAD_IMAGE}" + +# --- XLA Flags --- +XLA_FLAGS_ARRAY=( + "--xla_msa_enable_sync_slice_replacement=false" + "--xla_tpu_enable_sparse_core_collective_offload_2d_all_gather=true" + "--xla_tpu_enable_sparse_core_collective_offload_reduce_scatter=true" + "--xla_msa_enable_sync_copy_replacement=false" + "--xla_tpu_scoped_vmem_limit_kib=81000" + "--xla_tpu_enable_sparse_core_collective_offload_all_gather=true" + "--xla_tpu_enable_sparse_core_collective_offload_all_reduce=true" + "--xla_tpu_enable_concurrent_sparse_core_offloading=true" + "--xla_tpu_enable_sparse_core_offload_queuing_in_lhs=true" + "--xla_tpu_enable_layer_scheduler_for_dependent_collectives=true" + "--xla_tpu_use_single_sparse_core_for_all_gather_offload=true" + "--xla_tpu_sparse_core_all_gather_latency_multiplier=1" + "--xla_tpu_sparse_core_reduce_scatter_latency_multiplier=3" + "--xla_tpu_offload_gather_to_sparsecore=true" + "--xla_tpu_dvfs_p_state=7" + "--xla_tpu_disable_sparse_core_collective_offload_remover=true" + "--xla_tpu_use_tc_device_shape_on_sc=true" + "--xla_sc_enable_instruction_fusion=false" + "--xla_sc_disable_megacore_partitioning=true" + "--xla_tpu_enable_async_collective_fusion=true" + "--xla_tpu_overlap_compute_collective_tc=true" + "--xla_tpu_enable_async_collective_fusion_multiple_steps=true" + "--xla_tpu_enable_async_collective_fusion_fuse_all_gather=false" + "--xla_tpu_enable_async_collective_fusion_fuse_reduce_scatter=false" + "--xla_tpu_enable_async_collective_fusion_fuse_all_reduce=false" + "--xla_tpu_enable_latency_hiding_scheduler=true" + "--xla_latency_hiding_scheduler_rerun=10" + "--xla_tpu_all_gather_collective_matmul_mode=post_spmd_conservative" + "--xla_tpu_reduce_scatter_collective_matmul_mode=post_spmd_conservative" + "--xla_latency_hiding_scheduler_enable_selective_resources=true" + "--xla_tpu_enable_ilp_latency_hiding_scheduler=true" + "--xla_tpu_enable_all_experimental_scheduler_features=true" + "--xla_tpu_enable_scheduler_memory_pressure_tracking=true" + "--xla_tpu_host_transfer_overlap_limit=4" + "--xla_tpu_aggressive_opt_barrier_removal=ENABLED" + "--xla_lhs_prioritize_async_depth_over_stall=ENABLED" + "--xla_tpu_enable_ag_backward_pipelining=true" + "--xla_should_allow_loop_variant_parameter_in_chain=ENABLED" + "--xla_should_add_loop_invariant_op_in_chain=ENABLED" + "--xla_max_concurrent_host_send_recv=100" + "--xla_tpu_scheduler_percent_shared_memory_limit=150" +) +export XLA_FLAGS="${XLA_FLAGS_ARRAY[*]}" + +# --- MaxText Workload Overrides --- +MAXTEXT_ARGS_ARRAY=( + "model_name=${MODEL_NAME}" + "base_output_directory=${BASE_OUTPUT_DIR}" + "run_name=param-3" + "dataset_type=synthetic" + "dataset_name=synthetic" + "dtype=bfloat16" + "allow_split_physical_axes=True" + "ici_expert_parallelism=4" + "use_ring_of_experts=True" + "custom_mesh=hybrid_ring_64x4" + "use_ragged_sort=True" + "use_random_routing=True" + "num_moe_token_chunks=2" + "per_device_batch_size=8" + "opt_type=adamw" + "max_target_length=2048" + "ragged_buffer_factor=1.5" + "remat_policy=custom" + "reuse_example_batch=1" + "decoder_layer_input=device" + "ici_fsdp_parallelism=-1" + "steps=20" + "sa_block_q=1024" + "sa_block_kv=1024" + "sa_block_kv_compute=512" + "sa_block_q_dkv=1024" + "sa_block_kv_dkv=1024" + "sa_block_kv_dkv_compute=1024" + "sa_fuse_reciprocal=false" + "use_splash_scheduler=true" + "sa_use_base2_exp=true" + "dq_reduction_steps=3" + "hardware=tpu" + "skip_jax_distributed_system=False" + "attention=flash" + "use_tokamax_splash=True" + "sa_use_fused_bwd_kernel=True" + "sparse_matmul=True" + "megablox=True" + "wi_tile_fwd_batch_seq=128" + "wi_tile_dlhs_batch_seq=128" + "wi_tile_drhs_batch_seq=128" + "wo_tile_fwd_batch_seq=128" + "wo_tile_dlhs_batch_seq=128" + "wo_tile_drhs_batch_seq=128" + "wi_tile_fwd_embed_dim=3072" + "wi_tile_fwd_mlp_dim=1536" + "wi_tile_dlhs_embed_dim=3072" + "wi_tile_dlhs_mlp_dim=1536" + "wi_tile_drhs_embed_dim=3072" + "wi_tile_drhs_mlp_dim=1536" + "wo_tile_fwd_embed_dim=3072" + "wo_tile_fwd_mlp_dim=1536" + "wo_tile_dlhs_embed_dim=3072" + "wo_tile_dlhs_mlp_dim=1536" + "wo_tile_drhs_embed_dim=3072" + "wo_tile_drhs_mlp_dim=1536" + "use_tokamax_gmm=True" + "use_gmm_v2=True" + "optimizer_memory_host_offload=False" + "parameter_memory_host_offload=False" + "enable_checkpointing=False" + "async_checkpointing=False" + "tokenizer_type=tiktoken" + "tokenizer_path=tokenizer_74B/" + "override_model_config=true" + "mhc_expansion_rate=4" + "use_gdn_kernel=True" + "use_hybrid_gdn=True" + "profiler=xplane" + "profiler_steps=2" + "skip_first_n_steps_for_profiler=2" + "enable_tpu_profiling_options=True" + "upload_all_profiler_results=False" +) +MAXTEXT_ARGS="${MAXTEXT_ARGS_ARRAY[*]}" + +# The command to run inside the container +RUN_COMMAND="set -e && \ +export LIBTPU_INIT_ARGS=\"${XLA_FLAGS}\" && \ +export JAX_PLATFORMS='tpu,cpu' && \ +export ENABLE_PJRT_COMPATIBILITY='true' && \ +export JAX_DISTRIBUTED_INITIALIZE_TIMEOUT=1800 && \ +export PYTHONPATH=/deps:/deps/src:/deps/src/maxtext/src && \ +python3 src/maxtext/trainers/pre_train/train.py src/maxtext/configs/base.yml ${MAXTEXT_ARGS}" + +# --- XPK Workload Creation --- +echo "Creating XPK workload: ${WORKLOAD_NAME} on cluster: ${CLUSTER_NAME}" + +PYTHONPATH=/usr/local/google/home/chengnuojin/xpk/src python3 -m xpk.main workload create \ + --cluster="${CLUSTER_NAME}" \ + --project="${PROJECT_ID}" \ + --zone="${ZONE}" \ + --priority="${PRIORITY}" \ + --max-restarts="${MAX_RESTARTS}" \ + --device-type="${DEVICE_TYPE}" \ + --num-slices="${NUM_SLICES}" \ + --docker-image="${WORKLOAD_IMAGE}" \ + --workload="${WORKLOAD_NAME}" \ + --command="${RUN_COMMAND}" + +LOGS_URL="https://console.cloud.google.com/logs/query;query=resource.type%3D%22k8s_container%22%0Aresource.labels.project_id%3D%22${PROJECT_ID}%22%0Aresource.labels.location%3D%22southamerica-west1%22%0Aresource.labels.cluster_name%3D%22${CLUSTER_NAME}%22%0Aresource.labels.namespace_name%3D%22default%22%0Aresource.labels.pod_name%3A%22${WORKLOAD_NAME}-slice-job-0-0-%22%0Aseverity%3E%3DDEFAULT;storageScope=project;duration=P1D?project=${PROJECT_ID}" +GKE_URL="https://console.cloud.google.com/kubernetes/service/southamerica-west1/${CLUSTER_NAME}/default/${WORKLOAD_NAME}/details?project=${PROJECT_ID}" +TB_URL="https://tensorboard.corp.google.com/?logdir=${BASE_OUTPUT_DIR}/param-3/tensorboard" + +echo "========================================================================" +echo "πŸ“‹ Pantheon Cloud Logging (Worker 0 Logs):" +echo "${LOGS_URL}" +echo "" +echo "☸️ GKE Workload Details:" +echo "${GKE_URL}" +echo "" +echo "πŸ“Š GCS TensorBoard Link:" +echo "${TB_URL}" +echo "========================================================================" diff --git a/run_qwen3_80b_xpk_fsdp.sh b/run_qwen3_80b_xpk_fsdp.sh new file mode 100755 index 0000000000..9fc49351ca --- /dev/null +++ b/run_qwen3_80b_xpk_fsdp.sh @@ -0,0 +1,196 @@ +#!/bin/bash +set -e + +# Activate Python virtual environment +source /usr/local/google/home/chengnuojin/.venv/bin/activate + +# --- Environment Variables --- +export PROJECT_ID="tpu-prod-env-one-vm" +export CLUSTER_NAME="bodaborg-v6e-256-lcscld-c" +export ZONE="southamerica-west1-a" + +# --- Configuration & Automated Image Build --- +TIMESTAMP=$(date +%m%d%H%M%S) +export WORKLOAD_IMAGE="gcr.io/tpu-prod-env-one-vm/param3_21jul:chengnuojin_${TIMESTAMP}" +export WORKLOAD_NAME="chengnuojin-qn80b-fsdp-${TIMESTAMP}" +export DEVICE_TYPE="v6e-256" +export NUM_SLICES=1 +export PRIORITY="very-high" +export MAX_RESTARTS=1 +export NUM_STEPS=20 +export MODEL_NAME="qwen3-next-80b-a3b" +export BASE_OUTPUT_DIR="gs://chengnuojin-maxtext-logs/qwen3-next-80b-profiles/run-${TIMESTAMP}" + +echo "========================================================================" +echo "Building and uploading Docker runner image from /usr/local/google/home/chengnuojin/maxtext" +echo "Target Image: ${WORKLOAD_IMAGE}" +echo "========================================================================" + +( + cd /usr/local/google/home/chengnuojin/maxtext && \ + if ! docker image inspect maxtext_base_image &> /dev/null; then + echo "Local image 'maxtext_base_image' not found. Pulling gcr.io/tpu-prod-env-one-vm/param3_21jul:latest and tagging as maxtext_base_image..." + docker pull gcr.io/tpu-prod-env-one-vm/param3_21jul:latest + docker tag gcr.io/tpu-prod-env-one-vm/param3_21jul:latest maxtext_base_image + fi && \ + CLOUD_IMAGE_NAME="${WORKLOAD_IMAGE}" \ + BASE_IMAGE="gcr.io/tpu-prod-env-one-vm/param3_21jul:latest" \ + bash src/dependencies/scripts/docker_upload_runner.sh +) + +echo "Docker image upload complete: ${WORKLOAD_IMAGE}" + +# --- XLA Flags --- +XLA_FLAGS_ARRAY=( +"--xla_msa_enable_sync_slice_replacement=false" +"--xla_tpu_enable_sparse_core_collective_offload_2d_all_gather=true" +"--xla_tpu_enable_sparse_core_collective_offload_reduce_scatter=true" +"--xla_msa_enable_sync_copy_replacement=false" +"--xla_tpu_scoped_vmem_limit_kib=81000" +"--xla_tpu_enable_sparse_core_collective_offload_all_gather=true" +"--xla_tpu_enable_sparse_core_collective_offload_all_reduce=true" +"--xla_tpu_enable_concurrent_sparse_core_offloading=true" +"--xla_tpu_enable_sparse_core_offload_queuing_in_lhs=true" +"--xla_tpu_enable_layer_scheduler_for_dependent_collectives=true" +"--xla_tpu_use_single_sparse_core_for_all_gather_offload=true" +"--xla_tpu_sparse_core_all_gather_latency_multiplier=1" +"--xla_tpu_sparse_core_reduce_scatter_latency_multiplier=3" +"--xla_tpu_offload_gather_to_sparsecore=true" +"--xla_tpu_dvfs_p_state=7" +"--xla_tpu_disable_sparse_core_collective_offload_remover=true" +"--xla_tpu_use_tc_device_shape_on_sc=true" +"--xla_sc_enable_instruction_fusion=false" +"--xla_sc_disable_megacore_partitioning=true" +"--xla_tpu_enable_async_collective_fusion=true" +"--xla_tpu_overlap_compute_collective_tc=true" +"--xla_tpu_enable_async_collective_fusion_multiple_steps=true" +"--xla_tpu_enable_async_collective_fusion_fuse_all_gather=false" +"--xla_tpu_enable_async_collective_fusion_fuse_reduce_scatter=false" +"--xla_tpu_enable_async_collective_fusion_fuse_all_reduce=false" +"--xla_tpu_enable_latency_hiding_scheduler=true" +"--xla_latency_hiding_scheduler_rerun=10" +"--xla_tpu_all_gather_collective_matmul_mode=post_spmd_conservative" +"--xla_tpu_reduce_scatter_collective_matmul_mode=post_spmd_conservative" +"--xla_latency_hiding_scheduler_enable_selective_resources=true" +"--xla_tpu_enable_ilp_latency_hiding_scheduler=true" +"--xla_tpu_enable_all_experimental_scheduler_features=true" +"--xla_tpu_enable_scheduler_memory_pressure_tracking=true" +"--xla_tpu_host_transfer_overlap_limit=4" +"--xla_tpu_aggressive_opt_barrier_removal=ENABLED" +"--xla_lhs_prioritize_async_depth_over_stall=ENABLED" +"--xla_tpu_enable_ag_backward_pipelining=true" +"--xla_should_allow_loop_variant_parameter_in_chain=ENABLED" +"--xla_should_add_loop_invariant_op_in_chain=ENABLED" +"--xla_max_concurrent_host_send_recv=100" +"--xla_tpu_scheduler_percent_shared_memory_limit=150" +) +export XLA_FLAGS="${XLA_FLAGS_ARRAY[*]}" + +# --- MaxText Workload Overrides --- +MAXTEXT_ARGS_ARRAY=( + "model_name=${MODEL_NAME}" + "base_output_directory=${BASE_OUTPUT_DIR}" + "run_name=param-3" + "dataset_type=synthetic" + "dataset_name=synthetic" + "dtype=bfloat16" + "allow_split_physical_axes=False" + "ici_expert_parallelism=1" + "per_device_batch_size=6" + "opt_type=adamw" + "max_target_length=2048" + "remat_policy=custom" + "reuse_example_batch=1" + "decoder_layer_input=device" + "ici_fsdp_parallelism=-1" + "steps=20" + "sa_q_layout=SEQ_MINOR" + "sa_k_layout=HEAD_DIM_MINOR" + "sa_v_layout=HEAD_DIM_MINOR" + "sa_block_q=2048" + "sa_block_kv=2048" + "sa_block_kv_compute=1024" + "sa_block_q_dkv=2048" + "sa_block_kv_dkv=2048" + "sa_block_kv_dkv_compute=1024" + "hardware=tpu" + "skip_jax_distributed_system=False" + "attention=flash" + "use_tokamax_splash=True" + "sa_use_fused_bwd_kernel=True" + "sparse_matmul=True" + "megablox=True" + "wi_tile_fwd_batch_seq=64" + "wi_tile_dlhs_batch_seq=64" + "wi_tile_drhs_batch_seq=64" + "wo_tile_fwd_batch_seq=64" + "wo_tile_dlhs_batch_seq=64" + "wo_tile_drhs_batch_seq=64" + "wi_tile_fwd_embed_dim=3072" + "wi_tile_fwd_mlp_dim=1536" + "wi_tile_dlhs_embed_dim=3072" + "wi_tile_dlhs_mlp_dim=1536" + "wi_tile_drhs_embed_dim=3072" + "wi_tile_drhs_mlp_dim=1536" + "wo_tile_fwd_embed_dim=3072" + "wo_tile_fwd_mlp_dim=1536" + "wo_tile_dlhs_embed_dim=3072" + "wo_tile_dlhs_mlp_dim=1536" + "wo_tile_drhs_embed_dim=3072" + "wo_tile_drhs_mlp_dim=1536" + "use_tokamax_gmm=True" + "use_gmm_v2=True" + "optimizer_memory_host_offload=False" + "parameter_memory_host_offload=False" + "enable_checkpointing=False" + "async_checkpointing=False" + "tokenizer_type=tiktoken" + "tokenizer_path=tokenizer_74B/" + "use_gdn_kernel=True" + "use_hybrid_gdn=True" + "profiler=xplane" + "profiler_steps=2" + "skip_first_n_steps_for_profiler=2" + "enable_tpu_profiling_options=True" + "upload_all_profiler_results=False" +) +MAXTEXT_ARGS="${MAXTEXT_ARGS_ARRAY[*]}" + +# The command to run inside the container +RUN_COMMAND="set -e && \ +export LIBTPU_INIT_ARGS=\"${XLA_FLAGS}\" && \ +export JAX_PLATFORMS='tpu,cpu' && \ +export ENABLE_PJRT_COMPATIBILITY='true' && \ +export JAX_DISTRIBUTED_INITIALIZE_TIMEOUT=1800 && \ +export PYTHONPATH=/deps:/deps/src:/deps/src/maxtext/src && \ +python3 src/maxtext/trainers/pre_train/train.py src/maxtext/configs/base.yml ${MAXTEXT_ARGS}" + +# --- XPK Workload Creation --- +echo "Creating XPK workload: ${WORKLOAD_NAME} on cluster: ${CLUSTER_NAME}" + +PYTHONPATH=/usr/local/google/home/chengnuojin/xpk/src python3 -m xpk.main workload create \ + --cluster="${CLUSTER_NAME}" \ + --project="${PROJECT_ID}" \ + --zone="${ZONE}" \ + --priority="${PRIORITY}" \ + --max-restarts="${MAX_RESTARTS}" \ + --device-type="${DEVICE_TYPE}" \ + --num-slices="${NUM_SLICES}" \ + --docker-image="${WORKLOAD_IMAGE}" \ + --workload="${WORKLOAD_NAME}" \ + --command="${RUN_COMMAND}" + +LOGS_URL="https://console.cloud.google.com/logs/query;query=resource.type%3D%22k8s_container%22%0Aresource.labels.project_id%3D%22${PROJECT_ID}%22%0Aresource.labels.location%3D%22southamerica-west1%22%0Aresource.labels.cluster_name%3D%22${CLUSTER_NAME}%22%0Aresource.labels.namespace_name%3D%22default%22%0Aresource.labels.pod_name%3A%22${WORKLOAD_NAME}-slice-job-0-0-%22%0Aseverity%3E%3DDEFAULT;storageScope=project;duration=P1D?project=${PROJECT_ID}" +GKE_URL="https://console.cloud.google.com/kubernetes/service/southamerica-west1/${CLUSTER_NAME}/default/${WORKLOAD_NAME}/details?project=${PROJECT_ID}" +TB_URL="https://tensorboard.corp.google.com/?logdir=${BASE_OUTPUT_DIR}/param-3/tensorboard" + +echo "========================================================================" +echo "πŸ“‹ Pantheon Cloud Logging (Worker 0 Logs):" +echo "${LOGS_URL}" +echo "" +echo "☸️ GKE Workload Details:" +echo "${GKE_URL}" +echo "" +echo "πŸ“Š GCS TensorBoard Link:" +echo "${TB_URL}" +echo "========================================================================" diff --git a/src/dependencies/dockerfiles/maxtext_runner.Dockerfile b/src/dependencies/dockerfiles/maxtext_runner.Dockerfile index d85511c848..02b138931d 100644 --- a/src/dependencies/dockerfiles/maxtext_runner.Dockerfile +++ b/src/dependencies/dockerfiles/maxtext_runner.Dockerfile @@ -14,6 +14,9 @@ ENV MAXTEXT_REPO_ROOT=/deps # Set the working directory in the container WORKDIR /deps +# Install GDN v3 Tokamax commit +RUN pip install --no-deps --no-cache-dir --force-reinstall git+https://github.com/openxla/tokamax.git@b626dd8b54d708047788cf2ec538cba63a4e3739 + # Copy assets separately COPY ${PACKAGE_DIR}/maxtext/assets/ "${MAXTEXT_ASSETS_ROOT}" diff --git a/src/maxtext/checkpoint_conversion/utils/param_mapping.py b/src/maxtext/checkpoint_conversion/utils/param_mapping.py index 26359cdddc..a6f9dcdc0a 100644 --- a/src/maxtext/checkpoint_conversion/utils/param_mapping.py +++ b/src/maxtext/checkpoint_conversion/utils/param_mapping.py @@ -1386,100 +1386,225 @@ def QWEN3_NEXT_MAXTEXT_TO_HF_PARAM_MAPPING(config, maxtext_config, scan_layers=F } if scan_layers: - # 2. Scan over block cycles - for block_idx in range(layer_cycle_interval): - hf_indices = list(range(block_idx, num_main_layers, layer_cycle_interval)) - prefix = f"params-decoder-layers-layer_{block_idx}" + num_blocks = num_main_layers // layer_cycle_interval + num_scanned = num_blocks * layer_cycle_interval + num_remaining = num_main_layers % layer_cycle_interval - # Layer norms - mapping[f"{prefix}-input_layernorm-scale"] = [ # pyrefly: ignore[bad-assignment] - f"model.layers.{i}.input_layernorm.weight" for i in hf_indices - ] # pyrefly: ignore[bad-assignment] - mapping[f"{prefix}-post_attention_layernorm-scale"] = [ # pyrefly: ignore[bad-assignment] - f"model.layers.{i}.post_attention_layernorm.weight" for i in hf_indices - ] + def hf_layer(idx, suffix): + return f"model.layers.{idx}.{suffix}" - # Handle Interleaved Attention (Linear vs Full) - is_full_attention_layer = (block_idx + 1) % layer_cycle_interval == 0 + local_prefix = "params-decoder-scanned_blocks-local_layers" + local_positions = list(range(layer_cycle_interval - 1)) - if is_full_attention_layer: - mapping.update( # pyrefly: ignore[no-matching-overload] - { - f"{prefix}-attention-attention-query-kernel": [ - f"model.layers.{i}.self_attn.q_proj.weight" for i in hf_indices - ], - f"{prefix}-attention-attention-key-kernel": [ - f"model.layers.{i}.self_attn.k_proj.weight" for i in hf_indices - ], - f"{prefix}-attention-attention-value-kernel": [ - f"model.layers.{i}.self_attn.v_proj.weight" for i in hf_indices - ], - f"{prefix}-attention-attention-out-kernel": [ - f"model.layers.{i}.self_attn.o_proj.weight" for i in hf_indices - ], - f"{prefix}-attention-attention-query_norm-scale": [ - f"model.layers.{i}.self_attn.q_norm.weight" for i in hf_indices - ], - f"{prefix}-attention-attention-key_norm-scale": [ - f"model.layers.{i}.self_attn.k_norm.weight" for i in hf_indices - ], - } + # Local / linear attention layers (nested [block][local]) + mapping.update( + { + f"{local_prefix}-input_layernorm-scale": [ + [hf_layer(b * layer_cycle_interval + l, "input_layernorm.weight") for l in local_positions] + for b in range(num_blocks) + ], + f"{local_prefix}-post_attention_layernorm-scale": [ + [hf_layer(b * layer_cycle_interval + l, "post_attention_layernorm.weight") for l in local_positions] + for b in range(num_blocks) + ], + f"{local_prefix}-attention-in_proj_qkvz-kernel": [ + [hf_layer(b * layer_cycle_interval + l, "linear_attn.in_proj_qkvz.weight") for l in local_positions] + for b in range(num_blocks) + ], + f"{local_prefix}-attention-in_proj_ba-kernel": [ + [hf_layer(b * layer_cycle_interval + l, "linear_attn.in_proj_ba.weight") for l in local_positions] + for b in range(num_blocks) + ], + f"{local_prefix}-attention-conv1d-kernel": [ + [hf_layer(b * layer_cycle_interval + l, "linear_attn.conv1d.weight") for l in local_positions] + for b in range(num_blocks) + ], + f"{local_prefix}-attention-A_log": [ + [hf_layer(b * layer_cycle_interval + l, "linear_attn.A_log") for l in local_positions] + for b in range(num_blocks) + ], + f"{local_prefix}-attention-dt_bias": [ + [hf_layer(b * layer_cycle_interval + l, "linear_attn.dt_bias") for l in local_positions] + for b in range(num_blocks) + ], + f"{local_prefix}-attention-norm-rms_norm-scale": [ + [hf_layer(b * layer_cycle_interval + l, "linear_attn.norm.weight") for l in local_positions] + for b in range(num_blocks) + ], + f"{local_prefix}-attention-out_proj-kernel": [ + [hf_layer(b * layer_cycle_interval + l, "linear_attn.out_proj.weight") for l in local_positions] + for b in range(num_blocks) + ], + f"{local_prefix}-mlp-routed_experts-gate-kernel": [ + [hf_layer(b * layer_cycle_interval + l, "mlp.gate.weight") for l in local_positions] + for b in range(num_blocks) + ], + f"{local_prefix}-mlp-shared_expert-wi_0-kernel": [ + [hf_layer(b * layer_cycle_interval + l, "mlp.shared_expert.gate_proj.weight") for l in local_positions] + for b in range(num_blocks) + ], + f"{local_prefix}-mlp-shared_expert-wi_1-kernel": [ + [hf_layer(b * layer_cycle_interval + l, "mlp.shared_expert.up_proj.weight") for l in local_positions] + for b in range(num_blocks) + ], + f"{local_prefix}-mlp-shared_expert-wo-kernel": [ + [hf_layer(b * layer_cycle_interval + l, "mlp.shared_expert.down_proj.weight") for l in local_positions] + for b in range(num_blocks) + ], + f"{local_prefix}-mlp-shared_expert_gate-kernel": [ + [hf_layer(b * layer_cycle_interval + l, "mlp.shared_expert_gate.weight") for l in local_positions] + for b in range(num_blocks) + ], + f"{local_prefix}-mlp-routed_experts-wi_0": [ + [ + [hf_layer(b * layer_cycle_interval + l, f"mlp.experts.{e}.gate_proj.weight") for l in local_positions] + for b in range(num_blocks) + ] + for e in range(num_experts) + ], + f"{local_prefix}-mlp-routed_experts-wi_1": [ + [ + [hf_layer(b * layer_cycle_interval + l, f"mlp.experts.{e}.up_proj.weight") for l in local_positions] + for b in range(num_blocks) + ] + for e in range(num_experts) + ], + f"{local_prefix}-mlp-routed_experts-wo": [ + [ + [hf_layer(b * layer_cycle_interval + l, f"mlp.experts.{e}.down_proj.weight") for l in local_positions] + for b in range(num_blocks) + ] + for e in range(num_experts) + ], + } + ) + + global_prefix = "params-decoder-scanned_blocks-global_layer" + global_position = layer_cycle_interval - 1 + + # Global attention layer (flat over blocks) + mapping.update( + { + f"{global_prefix}-input_layernorm-scale": [ + hf_layer(b * layer_cycle_interval + global_position, "input_layernorm.weight") for b in range(num_blocks) + ], + f"{global_prefix}-post_attention_layernorm-scale": [ + hf_layer(b * layer_cycle_interval + global_position, "post_attention_layernorm.weight") + for b in range(num_blocks) + ], + f"{global_prefix}-attention-attention-query-kernel": [ + hf_layer(b * layer_cycle_interval + global_position, "self_attn.q_proj.weight") for b in range(num_blocks) + ], + f"{global_prefix}-attention-attention-key-kernel": [ + hf_layer(b * layer_cycle_interval + global_position, "self_attn.k_proj.weight") for b in range(num_blocks) + ], + f"{global_prefix}-attention-attention-value-kernel": [ + hf_layer(b * layer_cycle_interval + global_position, "self_attn.v_proj.weight") for b in range(num_blocks) + ], + f"{global_prefix}-attention-attention-out-kernel": [ + hf_layer(b * layer_cycle_interval + global_position, "self_attn.o_proj.weight") for b in range(num_blocks) + ], + f"{global_prefix}-attention-attention-query_norm-scale": [ + hf_layer(b * layer_cycle_interval + global_position, "self_attn.q_norm.weight") for b in range(num_blocks) + ], + f"{global_prefix}-attention-attention-key_norm-scale": [ + hf_layer(b * layer_cycle_interval + global_position, "self_attn.k_norm.weight") for b in range(num_blocks) + ], + f"{global_prefix}-mlp-routed_experts-gate-kernel": [ + hf_layer(b * layer_cycle_interval + global_position, "mlp.gate.weight") for b in range(num_blocks) + ], + f"{global_prefix}-mlp-shared_expert-wi_0-kernel": [ + hf_layer(b * layer_cycle_interval + global_position, "mlp.shared_expert.gate_proj.weight") + for b in range(num_blocks) + ], + f"{global_prefix}-mlp-shared_expert-wi_1-kernel": [ + hf_layer(b * layer_cycle_interval + global_position, "mlp.shared_expert.up_proj.weight") + for b in range(num_blocks) + ], + f"{global_prefix}-mlp-shared_expert-wo-kernel": [ + hf_layer(b * layer_cycle_interval + global_position, "mlp.shared_expert.down_proj.weight") + for b in range(num_blocks) + ], + f"{global_prefix}-mlp-shared_expert_gate-kernel": [ + hf_layer(b * layer_cycle_interval + global_position, "mlp.shared_expert_gate.weight") + for b in range(num_blocks) + ], + f"{global_prefix}-mlp-routed_experts-wi_0": [ + [ + hf_layer(b * layer_cycle_interval + global_position, f"mlp.experts.{e}.gate_proj.weight") + for b in range(num_blocks) + ] + for e in range(num_experts) + ], + f"{global_prefix}-mlp-routed_experts-wi_1": [ + [ + hf_layer(b * layer_cycle_interval + global_position, f"mlp.experts.{e}.up_proj.weight") + for b in range(num_blocks) + ] + for e in range(num_experts) + ], + f"{global_prefix}-mlp-routed_experts-wo": [ + [ + hf_layer(b * layer_cycle_interval + global_position, f"mlp.experts.{e}.down_proj.weight") + for b in range(num_blocks) + ] + for e in range(num_experts) + ], + } + ) + + # Remainder layers if any + if num_remaining > 0: + for rem_idx in range(num_remaining): + hf_layer_idx = num_scanned + rem_idx + prefix = f"params-decoder-layers_{hf_layer_idx}" + layer_in_block = rem_idx % layer_cycle_interval + is_full_attention_layer = (layer_in_block + 1) % layer_cycle_interval == 0 + mapping[f"{prefix}-input_layernorm-scale"] = f"model.layers.{hf_layer_idx}.input_layernorm.weight" + mapping[f"{prefix}-post_attention_layernorm-scale"] = ( + f"model.layers.{hf_layer_idx}.post_attention_layernorm.weight" ) - else: - # Linear/Hybrid Attention Block - mapping.update( # pyrefly: ignore[no-matching-overload] + if is_full_attention_layer: + mapping.update( + { + f"{prefix}-attention-attention-query-kernel": f"model.layers.{hf_layer_idx}.self_attn.q_proj.weight", + f"{prefix}-attention-attention-key-kernel": f"model.layers.{hf_layer_idx}.self_attn.k_proj.weight", + f"{prefix}-attention-attention-value-kernel": f"model.layers.{hf_layer_idx}.self_attn.v_proj.weight", + f"{prefix}-attention-attention-out-kernel": f"model.layers.{hf_layer_idx}.self_attn.o_proj.weight", + f"{prefix}-attention-attention-query_norm-scale": f"model.layers.{hf_layer_idx}.self_attn.q_norm.weight", + f"{prefix}-attention-attention-key_norm-scale": f"model.layers.{hf_layer_idx}.self_attn.k_norm.weight", + } + ) + else: + mapping.update( + { + f"{prefix}-attention-in_proj_qkvz-kernel": f"model.layers.{hf_layer_idx}.linear_attn.in_proj_qkvz.weight", + f"{prefix}-attention-in_proj_ba-kernel": f"model.layers.{hf_layer_idx}.linear_attn.in_proj_ba.weight", + f"{prefix}-attention-conv1d-kernel": f"model.layers.{hf_layer_idx}.linear_attn.conv1d.weight", + f"{prefix}-attention-A_log": f"model.layers.{hf_layer_idx}.linear_attn.A_log", + f"{prefix}-attention-dt_bias": f"model.layers.{hf_layer_idx}.linear_attn.dt_bias", + f"{prefix}-attention-norm-rms_norm-scale": f"model.layers.{hf_layer_idx}.linear_attn.norm.weight", + f"{prefix}-attention-out_proj-kernel": f"model.layers.{hf_layer_idx}.linear_attn.out_proj.weight", + } + ) + mapping.update( { - f"{prefix}-attention-in_proj_qkvz-kernel": [ - f"model.layers.{i}.linear_attn.in_proj_qkvz.weight" for i in hf_indices + f"{prefix}-mlp-routed_experts-gate-kernel": f"model.layers.{hf_layer_idx}.mlp.gate.weight", + f"{prefix}-mlp-shared_expert-wi_0-kernel": f"model.layers.{hf_layer_idx}.mlp.shared_expert.gate_proj.weight", + f"{prefix}-mlp-shared_expert-wi_1-kernel": f"model.layers.{hf_layer_idx}.mlp.shared_expert.up_proj.weight", + f"{prefix}-mlp-shared_expert-wo-kernel": f"model.layers.{hf_layer_idx}.mlp.shared_expert.down_proj.weight", + f"{prefix}-mlp-shared_expert_gate-kernel": f"model.layers.{hf_layer_idx}.mlp.shared_expert_gate.weight", + f"{prefix}-mlp-routed_experts-wi_0": [ + f"model.layers.{hf_layer_idx}.mlp.experts.{e}.gate_proj.weight" for e in range(num_experts) ], - f"{prefix}-attention-in_proj_ba-kernel": [ - f"model.layers.{i}.linear_attn.in_proj_ba.weight" for i in hf_indices + f"{prefix}-mlp-routed_experts-wi_1": [ + f"model.layers.{hf_layer_idx}.mlp.experts.{e}.up_proj.weight" for e in range(num_experts) ], - f"{prefix}-attention-conv1d-kernel": [f"model.layers.{i}.linear_attn.conv1d.weight" for i in hf_indices], - f"{prefix}-attention-A_log": [f"model.layers.{i}.linear_attn.A_log" for i in hf_indices], - f"{prefix}-attention-dt_bias": [f"model.layers.{i}.linear_attn.dt_bias" for i in hf_indices], - f"{prefix}-attention-norm-rms_norm-scale": [ - f"model.layers.{i}.linear_attn.norm.weight" for i in hf_indices - ], - f"{prefix}-attention-out_proj-kernel": [ - f"model.layers.{i}.linear_attn.out_proj.weight" for i in hf_indices + f"{prefix}-mlp-routed_experts-wo": [ + f"model.layers.{hf_layer_idx}.mlp.experts.{e}.down_proj.weight" for e in range(num_experts) ], } ) - - # 3. Handle MLP: Gates and Shared Experts - mapping.update( # pyrefly: ignore[no-matching-overload] - { - f"{prefix}-mlp-routed_experts-gate-kernel": [f"model.layers.{i}.mlp.gate.weight" for i in hf_indices], - f"{prefix}-mlp-shared_expert-wi_0-kernel": [ - f"model.layers.{i}.mlp.shared_expert.gate_proj.weight" for i in hf_indices - ], - f"{prefix}-mlp-shared_expert-wi_1-kernel": [ - f"model.layers.{i}.mlp.shared_expert.up_proj.weight" for i in hf_indices - ], - f"{prefix}-mlp-shared_expert-wo-kernel": [ - f"model.layers.{i}.mlp.shared_expert.down_proj.weight" for i in hf_indices - ], - f"{prefix}-mlp-shared_expert_gate-kernel": [ - f"model.layers.{i}.mlp.shared_expert_gate.weight" for i in hf_indices - ], - } - ) - - # 4. Handle MoE Routed Experts - mapping.update( # pyrefly: ignore[no-matching-overload] - { - f"{prefix}-mlp-routed_experts-wi_0": [ - [f"model.layers.{i}.mlp.experts.{e}.gate_proj.weight" for i in hf_indices] for e in range(num_experts) - ], - f"{prefix}-mlp-routed_experts-wi_1": [ - [f"model.layers.{i}.mlp.experts.{e}.up_proj.weight" for i in hf_indices] for e in range(num_experts) - ], - f"{prefix}-mlp-routed_experts-wo": [ - [f"model.layers.{i}.mlp.experts.{e}.down_proj.weight" for i in hf_indices] for e in range(num_experts) - ], - } - ) else: # Unscanned layer mapping for i in range(num_main_layers): @@ -1571,20 +1696,8 @@ def permute_conv(input_tensor, target_shape=None): "params-decoder-logits_dense-kernel": transpose, } - layer_cycle_interval = maxtext_config.inhomogeneous_layer_cycle_interval - num_main_layers = config["num_hidden_layers"] - loop_indices = range(layer_cycle_interval) if scan_layers else range(num_main_layers) - - for i in loop_indices: - if scan_layers: - prefix = f"params-decoder-layers-layer_{i}" - block_idx = i - else: - prefix = f"params-decoder-layers_{i}" - block_idx = i % layer_cycle_interval - is_full_attention_layer = (block_idx + 1) % layer_cycle_interval == 0 - - if is_full_attention_layer: + def _attach_block_hooks(prefix, is_global): + if is_global: for key in ["query", "key", "value", "out"]: hooks[f"{prefix}-attention-attention-{key}-kernel"] = reshape_kernel # pyrefly: ignore[bad-assignment] else: @@ -1604,6 +1717,16 @@ def permute_conv(input_tensor, target_shape=None): hooks[f"{mlp_prefix}-routed_experts-wi_1"] = transpose hooks[f"{mlp_prefix}-routed_experts-wo"] = transpose + if scan_layers: + _attach_block_hooks("params-decoder-scanned_blocks-local_layers", is_global=False) + _attach_block_hooks("params-decoder-scanned_blocks-global_layer", is_global=True) + else: + for i in range(config.base_num_decoder_layers): + prefix = f"params-decoder-layers_{i}" + block_idx = i % config.inhomogeneous_layer_cycle_interval + is_full_attention_layer = (block_idx + 1) % config.inhomogeneous_layer_cycle_interval == 0 + _attach_block_hooks(prefix, is_global=is_full_attention_layer) + return hooks diff --git a/src/maxtext/configs/base.yml b/src/maxtext/configs/base.yml index c21d5c5f7b..1684c10f94 100644 --- a/src/maxtext/configs/base.yml +++ b/src/maxtext/configs/base.yml @@ -1272,6 +1272,10 @@ gdn_num_value_heads: 32 gdn_chunk_size: 64 # Whether to apply L2 normalization to query and key tensors inside the Gated Delta Rule kernel. use_qk_norm_in_gdn: true +# Whether to use GDN Pallas kernel +use_gdn_kernel: false +# Whether to use hybrid GDN v3 Tokamax forward + Custom VJP backward +use_hybrid_gdn: false # The ratio of dimension to apply ROPE on partial_rotary_factor: 1.0 diff --git a/src/maxtext/configs/models/qwen3-next-80b-a3b.yml b/src/maxtext/configs/models/qwen3-next-80b-a3b.yml index 765977f1b5..327a395afb 100644 --- a/src/maxtext/configs/models/qwen3-next-80b-a3b.yml +++ b/src/maxtext/configs/models/qwen3-next-80b-a3b.yml @@ -18,35 +18,41 @@ decoder_block: "qwen3_next" # Core Architectural Parameters -base_emb_dim: 2048 -base_num_decoder_layers: 48 -base_num_query_heads: 16 -base_num_kv_heads: 2 -head_dim: 256 -vocab_size: 151936 +base_emb_dim: 3072 +base_num_decoder_layers: 40 +base_num_query_heads: 64 +base_num_kv_heads: 8 +head_dim: 128 +vocab_size: 128008 normalization_layer_epsilon: 1.0e-6 # MoE Specific Parameters -# Set base_mlp_dim to match base_moe_mlp_dim to pass validation for fully MoE models. -base_mlp_dim: 512 -base_moe_mlp_dim: 512 -num_experts: 512 +# base_mlp_dim sizes the dense-prefix layer's MLP +# base_moe_mlp_dim sizes every other (MoE) layer's routed + shared experts. +base_mlp_dim: 8192 +base_moe_mlp_dim: 1536 +num_experts: 128 shared_experts: 1 -num_experts_per_tok: 10 +num_experts_per_tok: 8 norm_topk_prob: true +# The first layer is a dense MLP (no MoE) and always uses full attention. +first_num_dense_layers: 1 + # Qwen3-Next Specific Parameters for Linear Attention (Gated Delta Net) -inhomogeneous_layer_cycle_interval: 4 +inhomogeneous_layer_cycle_interval: 3 gdn_conv_kernel_dim: 4 gdn_key_head_dim: 128 gdn_value_head_dim: 128 -gdn_num_key_heads: 16 -gdn_num_value_heads: 32 +gdn_num_key_heads: 32 +gdn_num_value_heads: 64 gdn_chunk_size: 64 # RoPE Settings -rope_max_timescale: 10000000 -partial_rotary_factor: 0.25 +rope_max_timescale: 10000 +partial_rotary_factor: 1.0 + +mhc_expansion_rate: 4 # General Model Settings enable_dropout: false diff --git a/src/maxtext/configs/types.py b/src/maxtext/configs/types.py index c4a270a567..5de50dce0e 100644 --- a/src/maxtext/configs/types.py +++ b/src/maxtext/configs/types.py @@ -1052,6 +1052,14 @@ class Qwen3Next(BaseModel): True, description="Whether to apply L2 normalization to query and key tensors inside the Gated Delta Rule kernel.", ) + use_gdn_kernel: bool = Field( + False, + description="Whether to use GDN Pallas kernel.", + ) + use_hybrid_gdn: bool = Field( + False, + description="Whether to use hybrid GDN v3 Tokamax forward + Custom VJP backward.", + ) partial_rotary_factor: float = Field(1.0, description="The ratio of dimension to apply ROPE on") @@ -1798,6 +1806,10 @@ class Muon(BaseModel): None, description="If None, apply width scaling to updates. If float, apply consistent rms scaling (recommend 0.2).", ) + muon_ns_steps: int = Field( + 5, + description="Number of Newton-Schulz iterations for Muon optimizer.", + ) class PositionalEmbedding(BaseModel): @@ -3778,6 +3790,7 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de DecoderBlockType.DEEPSEEK, DecoderBlockType.DEEPSEEK4, DecoderBlockType.QWEN3, + DecoderBlockType.QWEN3_NEXT, DecoderBlockType.GEMMA3, DecoderBlockType.LLAMA2, ]: diff --git a/src/maxtext/kernels/ragged/ragged_gather.py b/src/maxtext/kernels/ragged/ragged_gather.py index 0c2aabb1c1..bf2a7c78fb 100644 --- a/src/maxtext/kernels/ragged/ragged_gather.py +++ b/src/maxtext/kernels/ragged/ragged_gather.py @@ -410,7 +410,7 @@ def ragged_gather( # Guard against eager initialization on non-TPU hardware (e.g. during CPU tests). # pltpu.get_tpu_info() expects TPU hardware and will crash if executed on CPU. - if enforce_fallback or jax.devices()[0].platform != "tpu": + if enforce_fallback: return _fallback_implementation(x, indices, weights, has_weights) sc_info = pltpu.get_tpu_info().sparse_core diff --git a/src/maxtext/kernels/ragged/ragged_gather_reduce_v2.py b/src/maxtext/kernels/ragged/ragged_gather_reduce_v2.py index 0594676b74..a1cfb09d91 100644 --- a/src/maxtext/kernels/ragged/ragged_gather_reduce_v2.py +++ b/src/maxtext/kernels/ragged/ragged_gather_reduce_v2.py @@ -660,7 +660,7 @@ def ragged_gather_reduce( # Step 1: Choose the implementation (TensorCore fallback or SparseCore). # Guard against eager initialization on non-TPU hardware (e.g. during CPU tests). # pltpu.get_tpu_info() expects TPU hardware and will crash if executed on CPU. - if enforce_fallback or jax.devices()[0].platform != "tpu": + if enforce_fallback: return _fallback_implementation(x, indices, topk_weights, valid_rows_mask, reduce_group_size) sc_info = pltpu.get_tpu_info().sparse_core diff --git a/src/maxtext/layers/decoders.py b/src/maxtext/layers/decoders.py index 42753eb752..63f17facb3 100644 --- a/src/maxtext/layers/decoders.py +++ b/src/maxtext/layers/decoders.py @@ -877,7 +877,7 @@ def __call__( ) mhc_expand, mhc_reduce = mhc.get_functions(cfg.mhc_expansion_rate) - if cfg.mhc_expansion_rate > 1: + if cfg.mhc_expansion_rate > 1 and cfg.decoder_block in (DecoderBlockType.DEEPSEEK, DecoderBlockType.DEEPSEEK4): # (batch, length, emb_dim) --> (batch, length, mhc_expansion_rate, emb_dim) y = mhc_expand(y) @@ -1070,6 +1070,18 @@ def __call__( kv_caches=kv_caches, attention_metadata=attention_metadata, ) + elif cfg.decoder_block == DecoderBlockType.QWEN3_NEXT: + y = self._apply_qwen3_next_scanned_blocks( + y, + decoder_segment_ids, + decoder_positions, + deterministic, + model_mode, + previous_chunk, + slot, + kv_caches=kv_caches, + attention_metadata=attention_metadata, + ) elif cfg.decoder_block == DecoderBlockType.DEEPSEEK4: y = self._apply_deepseek4_scanned_blocks( y, @@ -1275,7 +1287,7 @@ def __call__( assert isinstance(y, jax.Array) # After the final transformer layer, `y` holds the raw, un-normalized hidden state. - if cfg.mhc_expansion_rate > 1: + if cfg.mhc_expansion_rate > 1 and cfg.decoder_block in (DecoderBlockType.DEEPSEEK, DecoderBlockType.DEEPSEEK4): if cfg.decoder_block == DecoderBlockType.DEEPSEEK4: hidden_state = mhc.DeepSeek4HyperHeadToLinen( config=cfg, @@ -1422,6 +1434,121 @@ def _apply_gemma3_scanned_blocks( return y + def _apply_qwen3_next_scanned_blocks( + self, + y, + decoder_segment_ids, + decoder_positions, + deterministic, + model_mode, + previous_chunk, + slot, + kv_caches=None, + attention_metadata=None, + ): + """Applies Qwen3-Next scanned decoder blocks, handling main scan and remainders.""" + + cfg = self.config + mesh = self.mesh + + # Define the repeating pattern length and calculate how many full blocks to scan + block_pattern_len = cfg.inhomogeneous_layer_cycle_interval + num_full_blocks = cfg.num_decoder_layers // block_pattern_len + remainder_layers = cfg.num_decoder_layers % block_pattern_len + + if num_full_blocks > 0: + ScannableBlockToLinen = qwen3.Qwen3NextScannableBlockToLinen + policy = self.get_remat_policy() + + kv_cache_scanned = maxtext_utils.prepare_kv_caches_for_scan( + kv_caches, num_full_blocks, block_pattern_len, stack=True + ) + + broadcast_args_spec = [ + (decoder_segment_ids, nn.broadcast), + (decoder_positions, nn.broadcast), + (deterministic, nn.broadcast), + (model_mode, nn.broadcast), + (slot, nn.broadcast), + (None, nn.broadcast), # page_state + (previous_chunk, nn.broadcast), + (None, nn.broadcast), # bidirectional_mask + (kv_cache_scanned, 0 if kv_caches is not None else nn.broadcast), + (attention_metadata, nn.broadcast), + ] + broadcast_args = tuple(arg for arg, _ in broadcast_args_spec) + in_axes_tuple = tuple(axis for _, axis in broadcast_args_spec) + + # For a fully scanned block, apply it inside an nn.scan over the calculated number of full blocks + y, returned_kv_cache = nn.scan( + ScannableBlockToLinen, + variable_axes={ + "params": cfg.param_scan_axis, + "cache": 0, + "intermediates": 0, + "aqt": 0, + "_overwrite_with_gradient": 0, + }, + split_rngs={"params": True, "dropout": cfg.enable_dropout}, + in_axes=in_axes_tuple, + length=num_full_blocks, + unroll=1, + metadata_params={ + nn.PARTITION_NAME: "layers", + "abstract_init": False, + }, + )( + config=cfg, + mesh=mesh, + quant=self.quant, + model_mode=model_mode, + num_of_layers=block_pattern_len, + remat_policy_fn=policy, + apply_internal_remat=True, + name="scanned_blocks", + )( + y, *broadcast_args + ) + + maxtext_utils.update_kv_caches_after_scan( + kv_caches, returned_kv_cache, num_full_blocks, block_pattern_len, stacked=True + ) + + # Process any remaining layers that don't fit into a full scanned block + for layer_id in range(cfg.num_decoder_layers - remainder_layers, cfg.num_decoder_layers): + layer = qwen3.Qwen3NextDecoderLayerToLinen( + config=cfg, + mesh=mesh, + model_mode=model_mode, + quant=self.quant, + layer_idx=layer_id, + ) + kv_cache = kv_caches[layer_id] if kv_caches is not None else None + + remainder_args = ( + decoder_segment_ids, + decoder_positions, + deterministic, + model_mode, + previous_chunk, + slot, + kv_cache, + attention_metadata, + ) + + y_and_kv = layer(y, *remainder_args) + if isinstance(y_and_kv, tuple): + y = y_and_kv[0] + new_kv = y_and_kv[1] + else: + y = y_and_kv + new_kv = None + + if kv_caches is not None and new_kv is not None: + kv_caches[layer_id] = new_kv + + return y + def _apply_gemma4_scanned_blocks( self, y, diff --git a/src/maxtext/layers/nnx_decoders.py b/src/maxtext/layers/nnx_decoders.py index 895ea27c14..963545c2c5 100644 --- a/src/maxtext/layers/nnx_decoders.py +++ b/src/maxtext/layers/nnx_decoders.py @@ -437,6 +437,7 @@ def __init__( self.is_gemma3 = self.config.decoder_block == DecoderBlockType.GEMMA3 self.is_gemma4 = self.config.decoder_block == DecoderBlockType.GEMMA4 self.is_gemma4_small = self.config.decoder_block == DecoderBlockType.GEMMA4_SMALL + self.is_qwen3_next = self.config.decoder_block == DecoderBlockType.QWEN3_NEXT if config.mhc_expansion_rate > 1 and config.decoder_block == DecoderBlockType.DEEPSEEK4: self.hc_head = mhc.DeepSeek4HyperHead( @@ -547,6 +548,8 @@ def _init_scanned_layers(self, decoder_block_classes, rngs, mesh): self._init_scanned_gemma3(decoder_block_classes, rngs, mesh) elif self.is_gemma4: self._init_scanned_gemma4(decoder_block_classes, rngs, mesh) + elif self.is_qwen3_next: + self._init_scanned_qwen3_next(decoder_block_classes, rngs, mesh) else: self._init_scanned_generic(decoder_block_classes, rngs) @@ -717,6 +720,43 @@ def _init_scanned_gemma4(self, decoder_block_classes, rngs, mesh): rngs=rngs, ) + def _init_scanned_qwen3_next(self, decoder_block_classes, rngs, mesh): + """Initializes scanned Qwen3-Next layers.""" + config = self.config + cycle_interval = config.inhomogeneous_layer_cycle_interval + scan_length = config.num_decoder_layers // cycle_interval + num_remaining_layers = config.num_decoder_layers % cycle_interval + policy = self.get_remat_policy() + layer_kwargs = { + "num_of_layers": cycle_interval, + "remat_policy_fn": policy, + "apply_internal_remat": True, + } + rem_layer_kwargs = { + "num_of_layers": num_remaining_layers, + "remat_policy_fn": policy, + "apply_internal_remat": True, + } + + RemattedQwen3NextBlock = qwen3.Qwen3NextScannableBlock + + if scan_length > 0: + self.scanned_blocks = self._create_scanned_layers( + RemattedQwen3NextBlock, + length=scan_length, + metadata_axis_name="layers", + rngs=rngs, + **layer_kwargs, + ) + self.layers_remainder = RemattedQwen3NextBlock( + config=self.config, + mesh=mesh, + quant=self.quant, + model_mode=self.model_mode, + **rem_layer_kwargs, + rngs=rngs, + ) + def _init_scanned_generic(self, decoder_block_classes, rngs): """Initializes scanned generic decoder layers.""" config = self.config @@ -1604,7 +1644,7 @@ def __call__( mhc_reduce = None if hasattr(cfg, "mhc_expansion_rate"): mhc_expand, mhc_reduce = mhc.get_functions(cfg.mhc_expansion_rate) - if cfg.mhc_expansion_rate > 1: + if cfg.mhc_expansion_rate > 1 and cfg.decoder_block in (DecoderBlockType.DEEPSEEK, DecoderBlockType.DEEPSEEK4): # (batch, length, emb_dim) --> (batch, length, mhc_expansion_rate, emb_dim) y = mhc_expand(y) @@ -1854,6 +1894,13 @@ def __call__( layer_kwargs, kv_caches=kv_caches, ) + elif self.is_qwen3_next: + y = self._apply_qwen3_next_scanned_blocks( + y, + layer_args, + layer_kwargs, + kv_caches=kv_caches, + ) else: scan_length = int(cfg.num_decoder_layers / cfg.inhomogeneous_layer_cycle_interval) if kv_caches is not None: @@ -1975,7 +2022,7 @@ def pure_layer_fn(graphdef_in, state_in, y_in, kv_in): assert isinstance(y, jax.Array) # After the final transformer layer, `y` holds the raw, un-normalized hidden state. - if getattr(cfg, "mhc_expansion_rate", 1) > 1: + if getattr(cfg, "mhc_expansion_rate", 1) > 1 and cfg.decoder_block in (DecoderBlockType.DEEPSEEK, DecoderBlockType.DEEPSEEK4): if cfg.decoder_block == DecoderBlockType.DEEPSEEK4: hidden_state = self.hc_head(y) else: @@ -2214,6 +2261,71 @@ def pure_gemma_fn(graphdef, state_in, y_in, kv_in): return y + def _apply_qwen3_next_scanned_blocks( + self, + y, + layer_args, + layer_kwargs, + kv_caches=None, + ): + """Applies Qwen3-Next scanned decoder blocks, handling main scan and remainders.""" + + cfg = self.config + cycle_interval = cfg.inhomogeneous_layer_cycle_interval + scan_length = cfg.num_decoder_layers // cycle_interval + + block_unroll = 1 + if scan_length > 0: + grouped_kv_caches = maxtext_utils.prepare_kv_caches_for_scan(kv_caches, scan_length, cycle_interval, stack=False) + y, self.scanned_blocks, _ = self._apply_layers_sequentially( + self.scanned_blocks, + y, + *layer_args, + length=scan_length, + kv_caches_stacked=grouped_kv_caches, + skip_block_remat=True, + unroll=block_unroll, + **layer_kwargs, + ) + maxtext_utils.update_kv_caches_after_scan(kv_caches, grouped_kv_caches, scan_length, cycle_interval, stacked=False) + + num_remaining_layers = cfg.num_decoder_layers % cycle_interval + if num_remaining_layers > 0: + policy = self.get_remat_policy() + prevent_cse = maxtext_utils.should_prevent_cse_in_remat(cfg) + + remainder_kv = None + if kv_caches is not None: + start_idx = scan_length * cycle_interval + remainder_kv = tuple(kv_caches[start_idx : start_idx + num_remaining_layers]) + + def pure_qwen3_fn(graphdef, state_in, y_in, kv_in): + merged_layer = nnx.merge(graphdef, state_in) + call_kwargs = dict(layer_kwargs) + if kv_in is not None: + call_kwargs["kv_cache"] = kv_in + out_res = merged_layer(y_in, *layer_args, **call_kwargs) + if isinstance(out_res, tuple): + out_y = out_res[0] + out_kv = out_res[1] if len(out_res) > 1 else None + else: + out_y = out_res + out_kv = None + return out_y, out_kv, nnx.state(merged_layer) + + checkpointed_qwen3_fn = jax.checkpoint(pure_qwen3_fn, policy=policy, prevent_cse=prevent_cse) + + graphdef, state = nnx.split(self.layers_remainder) + y, updated_remainder_kv, new_state = checkpointed_qwen3_fn(graphdef, state, y, remainder_kv) + nnx.update(self.layers_remainder, new_state) + + if kv_caches is not None and updated_remainder_kv is not None: + start_idx = scan_length * cycle_interval + for offset, updated_item in enumerate(updated_remainder_kv): + kv_caches[start_idx + offset] = updated_item + + return y + def _apply_gemma4_small_layers( self, y, diff --git a/src/maxtext/models/hybrid_gdn.py b/src/maxtext/models/hybrid_gdn.py new file mode 100644 index 0000000000..74fa38136a --- /dev/null +++ b/src/maxtext/models/hybrid_gdn.py @@ -0,0 +1,279 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Hybrid Gated Delta Net (GDN) implementations for MaxText using Tokamax GDN v3 forward + Custom VJP backward.""" + +import functools +from typing import Any, Optional, Tuple + +import jax +import jax.numpy as jnp + + +def pure_jax_fused_conv1d_gdn( + qkv: jax.Array, + b: jax.Array, + a: jax.Array, + conv_weight: jax.Array, + conv_bias: Optional[jax.Array], + a_log: jax.Array, + dt_bias: jax.Array, + conv_state: Optional[jax.Array], + recurrent_state: Optional[jax.Array], + *, + num_k_heads: int, + num_v_heads: int, + head_k_dim: int, + head_v_dim: int, + conv_kernel_size: int, + chunk_size: int, + use_qk_norm_in_gdn: bool, + compute_dtype: jnp.dtype, +) -> Tuple[jax.Array, Tuple[jax.Array, jax.Array]]: + """Pure-JAX composite of Conv1D + GDN used during backward pass autodiff.""" + from maxtext.models.qwen3 import jax_chunk_gated_delta_rule + batch, seq_len, _ = qkv.shape + key_dim = num_k_heads * head_k_dim + + # --- Step B: Pure JAX 1D Convolution --- + conv_input = jnp.pad(qkv, ((0, 0), (conv_kernel_size - 1, 0), (0, 0))) + conv_weight_cast = conv_weight.astype(qkv.dtype) + conv_out = jax.lax.conv_general_dilated( + lhs=conv_input, + rhs=conv_weight_cast, + window_strides=(1,), + padding="VALID", + dimension_numbers=("NWC", "WIO", "NWC"), + feature_group_count=qkv.shape[-1], + ) + if conv_bias is not None: + conv_out = conv_out + conv_bias.astype(qkv.dtype) + conv_out = conv_out[:, -seq_len:, :] + qkv_conv = jax.nn.silu(conv_out.astype(jnp.float32)).astype(compute_dtype) + + q_conv, k_conv, v_conv = jnp.split(qkv_conv, [key_dim, 2 * key_dim], axis=-1) + + # Reshape for GDN + query = q_conv.reshape(batch, seq_len, num_k_heads, head_k_dim) + key = k_conv.reshape(batch, seq_len, num_k_heads, head_k_dim) + value = v_conv.reshape(batch, seq_len, num_v_heads, head_v_dim) + + A_log_cast = jnp.asarray(a_log, dtype=compute_dtype) + dt_bias_cast = jnp.asarray(dt_bias, dtype=compute_dtype) + beta = jax.nn.sigmoid(b) + g = -jnp.exp(A_log_cast) * jax.nn.softplus(a + dt_bias_cast) + + if num_v_heads > num_k_heads and num_v_heads % num_k_heads == 0: + repeats = num_v_heads // num_k_heads + query = jnp.repeat(query, repeats, axis=2) + key = jnp.repeat(key, repeats, axis=2) + + core_attn_out, next_recurrent_state = jax_chunk_gated_delta_rule( + query=query, + key=key, + value=value, + g=g, + beta=beta, + chunk_size=chunk_size, + initial_state=recurrent_state, + use_qk_norm_in_gdn=use_qk_norm_in_gdn, + compute_dtype=compute_dtype, + ) + + next_conv_state = qkv[:, -(conv_kernel_size - 1):, :] if seq_len >= conv_kernel_size - 1 else jnp.zeros((batch, conv_kernel_size - 1, qkv.shape[-1]), dtype=qkv.dtype) + if next_recurrent_state is None: + next_recurrent_state = jnp.zeros((batch, num_v_heads, head_k_dim, head_v_dim), dtype=compute_dtype) + + return core_attn_out.astype(qkv.dtype), (next_conv_state.astype(qkv.dtype), next_recurrent_state.astype(qkv.dtype)) + + +def _run_tokamax_fused_fwd( + qkv: jax.Array, + b: jax.Array, + a: jax.Array, + conv_weight: jax.Array, + conv_bias: Optional[jax.Array], + a_log: jax.Array, + dt_bias: jax.Array, + conv_state: Optional[jax.Array], + recurrent_state: Optional[jax.Array], + *, + num_k_heads: int, + num_v_heads: int, + head_k_dim: int, + head_v_dim: int, + conv_kernel_size: int, + chunk_size: int, + use_qk_norm_in_gdn: bool, + compute_dtype: jnp.dtype, +): + if jax.default_backend() != "tpu": + return pure_jax_fused_conv1d_gdn( + qkv, b, a, conv_weight, conv_bias, a_log, dt_bias, conv_state, recurrent_state, + num_k_heads=num_k_heads, num_v_heads=num_v_heads, head_k_dim=head_k_dim, head_v_dim=head_v_dim, + conv_kernel_size=conv_kernel_size, chunk_size=chunk_size, use_qk_norm_in_gdn=use_qk_norm_in_gdn, compute_dtype=compute_dtype, + ) + + # When on TPU, invoke Tokamax GDN v3 fused_conv1d_gdn kernel + from tokamax._src.ops.experimental.causal_conv1d_gated_delta_rule import wrapper as tokamax_gdn_wrapper + batch_size, seq_len, dim_size = qkv.shape + num_seqs = batch_size + + qkv_flat = qkv.reshape(-1, dim_size) + b_flat = b.reshape(-1, b.shape[-1]) + a_flat = a.reshape(-1, a.shape[-1]) + tokamax_conv_weight = jnp.swapaxes(conv_weight, 0, 2) + + query_start_loc = jnp.arange(0, (num_seqs + 1) * seq_len, seq_len, dtype=jnp.int32) + state_indices = jnp.arange(num_seqs, dtype=jnp.int32) + seq_lens = jnp.full((num_seqs,), seq_len, dtype=jnp.int32) + distribution = jnp.array([0, 0, num_seqs], dtype=jnp.int32) + + if conv_state is None: + tokamax_conv_state = jnp.zeros((num_seqs + 1, conv_kernel_size - 1, dim_size), dtype=qkv.dtype) + elif conv_state.shape[0] == num_seqs: + tokamax_conv_state = jnp.pad(conv_state, ((1, 0), (0, 0), (0, 0))) + else: + tokamax_conv_state = conv_state + + if recurrent_state is None: + tokamax_recurrent_state = jnp.zeros((num_seqs + 1, num_v_heads, head_k_dim, head_v_dim), dtype=qkv.dtype) + elif recurrent_state.shape[0] == num_seqs: + tokamax_recurrent_state = jnp.pad(recurrent_state, ((1, 0), (0, 0), (0, 0), (0, 0))) + else: + tokamax_recurrent_state = recurrent_state + + (new_conv_state, new_recurrent_state), core_attn_out_flat = tokamax_gdn_wrapper.fused_conv1d_gdn( + qkv=qkv_flat, + b=b_flat, + a=a_flat, + conv_state=tokamax_conv_state, + recurrent_state=tokamax_recurrent_state, + conv_weight=tokamax_conv_weight, + conv_bias=conv_bias, + a_log=a_log, + dt_bias=dt_bias, + query_start_loc=query_start_loc, + state_indices=state_indices, + distribution=distribution, + seq_lens=seq_lens, + n_kq=num_k_heads, + n_v=num_v_heads, + d_k=head_k_dim, + d_v=head_v_dim, + kernel_size=conv_kernel_size, + compute_precision=jnp.dtype(jnp.float32), + ) + + core_attn_out = core_attn_out_flat.reshape(batch_size, seq_len, num_v_heads, head_v_dim) + return core_attn_out.astype(qkv.dtype), (new_conv_state[1:].astype(qkv.dtype), new_recurrent_state[1:].astype(qkv.dtype)) + + +@functools.partial(jax.custom_vjp, nondiff_argnums=(9, 10, 11, 12, 13, 14, 15, 16)) +def hybrid_fused_conv1d_gdn( + qkv: jax.Array, + b: jax.Array, + a: jax.Array, + conv_weight: jax.Array, + conv_bias: Optional[jax.Array], + a_log: jax.Array, + dt_bias: jax.Array, + conv_state: Optional[jax.Array], + recurrent_state: Optional[jax.Array], + num_k_heads: int, + num_v_heads: int, + head_k_dim: int, + head_v_dim: int, + conv_kernel_size: int, + chunk_size: int, + use_qk_norm_in_gdn: bool, + compute_dtype: jnp.dtype, +) -> Tuple[jax.Array, Tuple[jax.Array, jax.Array]]: + """Hybrid Fused Conv1D + GDN: Tokamax GDN v3 forward + Custom VJP backward.""" + return _run_tokamax_fused_fwd( + qkv, b, a, conv_weight, conv_bias, a_log, dt_bias, conv_state, recurrent_state, + num_k_heads=num_k_heads, num_v_heads=num_v_heads, head_k_dim=head_k_dim, head_v_dim=head_v_dim, + conv_kernel_size=conv_kernel_size, chunk_size=chunk_size, use_qk_norm_in_gdn=use_qk_norm_in_gdn, compute_dtype=compute_dtype, + ) + + +def _hybrid_fused_conv1d_gdn_fwd( + qkv: jax.Array, + b: jax.Array, + a: jax.Array, + conv_weight: jax.Array, + conv_bias: Optional[jax.Array], + a_log: jax.Array, + dt_bias: jax.Array, + conv_state: Optional[jax.Array], + recurrent_state: Optional[jax.Array], + num_k_heads: int, + num_v_heads: int, + head_k_dim: int, + head_v_dim: int, + conv_kernel_size: int, + chunk_size: int, + use_qk_norm_in_gdn: bool, + compute_dtype: jnp.dtype, +): + out, states = _run_tokamax_fused_fwd( + qkv, b, a, conv_weight, conv_bias, a_log, dt_bias, conv_state, recurrent_state, + num_k_heads=num_k_heads, num_v_heads=num_v_heads, head_k_dim=head_k_dim, head_v_dim=head_v_dim, + conv_kernel_size=conv_kernel_size, chunk_size=chunk_size, use_qk_norm_in_gdn=use_qk_norm_in_gdn, compute_dtype=compute_dtype, + ) + residuals = ( + qkv, b, a, conv_weight, conv_bias, a_log, dt_bias, conv_state, recurrent_state + ) + return (out, states), residuals + + +def _hybrid_fused_conv1d_gdn_bwd( + num_k_heads: int, + num_v_heads: int, + head_k_dim: int, + head_v_dim: int, + conv_kernel_size: int, + chunk_size: int, + use_qk_norm_in_gdn: bool, + compute_dtype: jnp.dtype, + residuals: tuple, + cotangents: tuple, +): + ( + qkv, b, a, conv_weight, conv_bias, a_log, dt_bias, conv_state, recurrent_state + ) = residuals + + def target_fn(qkv_, b_, a_, cw_, cb_, al_, dt_, cs_, rs_): + return pure_jax_fused_conv1d_gdn( + qkv_, b_, a_, cw_, cb_, al_, dt_, cs_, rs_, + num_k_heads=num_k_heads, + num_v_heads=num_v_heads, + head_k_dim=head_k_dim, + head_v_dim=head_v_dim, + conv_kernel_size=conv_kernel_size, + chunk_size=chunk_size, + use_qk_norm_in_gdn=use_qk_norm_in_gdn, + compute_dtype=compute_dtype, + ) + + _, vjp_fn = jax.vjp( + target_fn, + qkv, b, a, conv_weight, conv_bias, a_log, dt_bias, conv_state, recurrent_state, + ) + d_out, d_states = cotangents + d_conv_state, d_recurrent_state = d_states + return vjp_fn((d_out, (d_conv_state, d_recurrent_state))) + + +hybrid_fused_conv1d_gdn.defvjp(_hybrid_fused_conv1d_gdn_fwd, _hybrid_fused_conv1d_gdn_bwd) diff --git a/src/maxtext/models/qwen3.py b/src/maxtext/models/qwen3.py index 7cb710bf1e..be8dcffb7f 100644 --- a/src/maxtext/models/qwen3.py +++ b/src/maxtext/models/qwen3.py @@ -25,6 +25,7 @@ import jax.nn from jax import lax from jax.ad_checkpoint import checkpoint_name +from jax.experimental import xla_metadata from jax.sharding import Mesh import jax.numpy as jnp @@ -37,7 +38,9 @@ from maxtext.layers import attentions from maxtext.layers import initializers as max_initializers from maxtext.layers import moe -from maxtext.layers import nnx_wrappers +from maxtext.layers import mhc +from maxtext.common.common_types import HyperConnectionType +from maxtext.layers import nnx_scan, nnx_wrappers from maxtext.layers import quantizations from maxtext.layers.embeddings import Qwen3OmniMoeVisionPosEmbedInterpolate, PositionalEmbedding from maxtext.layers.normalizations import RMSNorm, l2norm, Qwen3NextRMSNorm, Qwen3NextRMSNormGated @@ -46,7 +49,7 @@ from maxtext.layers.linears import DenseGeneral, MlpBlock from maxtext.layers.moe import RoutedMoE from maxtext.layers.initializers import nd_dense_init, variable_to_logically_partitioned -from maxtext.utils import max_utils +from maxtext.utils import max_utils, maxtext_utils from maxtext.inference import kvcache @@ -170,7 +173,7 @@ def scan_body(prev_state, x): return new_last_recurrent_state, core_attn_out_i - final_state, core_attn_out_stacked = jax.lax.scan(scan_body, last_recurrent_state, xs) + final_state, core_attn_out_stacked = jax.lax.scan(scan_body, last_recurrent_state, xs, unroll=0) core_attn_out = jnp.transpose(core_attn_out_stacked, (1, 2, 0, 3, 4)) core_attn_out = core_attn_out.reshape(batch_size, num_heads, -1, v_head_dim) @@ -180,6 +183,57 @@ def scan_body(prev_state, x): return core_attn_out, final_state if output_final_state else None +@jax.custom_vjp +def invert_unit_lower_triangular_log_depth(S): + """ + Computes (I + S)^-1 for a strictly lower triangular matrix S + using log-depth Newton-Schulz iterations. + + This is highly optimized for TPUs/GPUs and replaces + jax.scipy.linalg.solve_triangular for chunkwise linear attention. + """ + chunk_size = S.shape[-1] + + # Ensure S is strictly lower triangular (zero out diagonal and upper half) + # This guarantees mathematical correctness and stability + S_strict = jnp.tril(S, k=-1) + + # Base identity matrix + identity = jnp.eye(chunk_size, dtype=S.dtype) + + # Initial approximation and error term + A = identity - S_strict + E = jnp.tril(S_strict @ S_strict, k=-1) + + # Log-depth Taylor series exact computation + steps = int(math.ceil(math.log2(chunk_size))) + for _ in range(steps - 1): + # Update inverse and error using batched matmuls + A = jnp.tril(A + A @ E) + E = jnp.tril(E @ E, k=-1) + + return A + + +@functools.partial(jax.named_call, name="invert_triangular_fwd") +def _invert_unit_lower_triangular_log_depth_fwd(S): + A = invert_unit_lower_triangular_log_depth(S) + return A, A + + +@functools.partial(jax.named_call, name="invert_triangular_bwd") +def _invert_unit_lower_triangular_log_depth_bwd(res, g): + A = res + grad_S = jnp.tril(-(A.mT @ g @ A.mT), k=-1) + return (grad_S,) + + +invert_unit_lower_triangular_log_depth.defvjp( + _invert_unit_lower_triangular_log_depth_fwd, _invert_unit_lower_triangular_log_depth_bwd +) + + +@functools.partial(jax.named_call, name="jax_chunked_delta_rule") def jax_chunk_gated_delta_rule( query: Array, key: Array, @@ -264,11 +318,11 @@ def to_chunk_scalar(x): S = S * jnp.exp(g_diff) S = jnp.where(mask, S, 0.0) - # Inversion (A) - Strictly float32 - identity = jnp.eye(chunk_size, dtype=jnp.float32) - identity_broadcasted = jnp.broadcast_to(identity, S.shape) + # Cast to float32 explicitly as you were doing before + S = S.astype(jnp.float32) - A = jax.scipy.linalg.solve_triangular(identity + S, identity_broadcasted, lower=True, unit_diagonal=True) + # Inversion (A) - Replaces solve_triangular entirely + A = invert_unit_lower_triangular_log_depth(S) # 5. WY Factors v_beta = v_c * beta_c[..., None] @@ -726,6 +780,7 @@ def __call__( # Reshape GDN output and apply gated norm + out projection. gdn_output = gdn_output.reshape(batch, seq_len, self.num_v_heads, self.head_v_dim) + gdn_output = checkpoint_name(gdn_output, "context") gated_output = self.norm(gdn_output, z) gated_output = gated_output.reshape(batch, seq_len, -1) output = self.out_proj(gated_output) @@ -839,6 +894,116 @@ def extract_state(c_in, v_len): use_qk_norm_in_gdn=cfg.use_qk_norm_in_gdn, compute_dtype=cfg.dtype, ) + elif getattr(cfg, "use_gdn_kernel", False) and getattr(cfg, "use_hybrid_gdn", False): + from maxtext.models.hybrid_gdn import hybrid_fused_conv1d_gdn + + if self.mesh is not None: + logical_rules = get_logical_axis_rules() + batch_pspec3 = logical_to_mesh_axes((KV_BATCH, None, None), mesh=self.mesh, rules=logical_rules) + batch_pspec4 = logical_to_mesh_axes((KV_BATCH, None, None, None), mesh=self.mesh, rules=logical_rules) + none_pspec3 = logical_to_mesh_axes((None, None, None), mesh=self.mesh, rules=logical_rules) + none_pspec1 = logical_to_mesh_axes((None,), mesh=self.mesh, rules=logical_rules) + + recurrent_state_arg = ( + recurrent_state + if recurrent_state is not None + else jnp.zeros((batch, self.num_v_heads, self.head_k_dim, self.head_v_dim), dtype=cfg.dtype) + ) + conv_state_arg = ( + conv_state + if conv_state is not None + else jnp.zeros((batch, self.config.gdn_conv_kernel_dim - 1, qkv.shape[-1]), dtype=cfg.dtype) + ) + conv_bias_arg = ( + self.conv1d.bias.value + if hasattr(self.conv1d, "bias") and self.conv1d.bias is not None + else jnp.zeros((qkv.shape[-1],), dtype=cfg.dtype) + ) + + @functools.partial( + jax.shard_map, + mesh=self.mesh, + in_specs=( + batch_pspec3, # qkv + batch_pspec3, # b + batch_pspec3, # a + none_pspec3, # conv_weight + none_pspec1, # conv_bias + none_pspec1, # a_log + none_pspec1, # dt_bias + batch_pspec3, # conv_state + batch_pspec4, # recurrent_state + ), + out_specs=( + batch_pspec4, # core_attn_out + (batch_pspec3, batch_pspec4), # (next_conv_state, next_recurrent_state) + ), + check_vma=False, + ) + def shard_mapped_hybrid_gdn(qkv_val, b_val, a_val, cw_val, cb_val, alog_val, dt_val, cs_val, rs_val): + return hybrid_fused_conv1d_gdn( + qkv=qkv_val, + b=b_val, + a=a_val, + conv_weight=cw_val, + conv_bias=cb_val, + a_log=alog_val, + dt_bias=dt_val, + conv_state=cs_val, + recurrent_state=rs_val, + num_k_heads=self.num_k_heads, + num_v_heads=self.num_v_heads, + head_k_dim=self.head_k_dim, + head_v_dim=self.head_v_dim, + conv_kernel_size=self.config.gdn_conv_kernel_dim, + chunk_size=self.config.gdn_chunk_size, + use_qk_norm_in_gdn=self.config.use_qk_norm_in_gdn, + compute_dtype=self.config.dtype, + ) + + core_attn_out, (next_conv_state, next_recurrent_state) = shard_mapped_hybrid_gdn( + qkv, + b, + a, + self.conv1d.kernel.value, + conv_bias_arg, + self.A_log[...], + self.dt_bias[...], + conv_state_arg, + recurrent_state_arg, + ) + else: + core_attn_out, (next_conv_state, next_recurrent_state) = hybrid_fused_conv1d_gdn( + qkv=qkv, + b=b, + a=a, + conv_weight=self.conv1d.kernel.value, + conv_bias=None, + a_log=self.A_log[...], + dt_bias=self.dt_bias[...], + conv_state=conv_state, + recurrent_state=recurrent_state, + num_k_heads=self.num_k_heads, + num_v_heads=self.num_v_heads, + head_k_dim=self.head_k_dim, + head_v_dim=self.head_v_dim, + conv_kernel_size=self.config.gdn_conv_kernel_dim, + chunk_size=self.config.gdn_chunk_size, + use_qk_norm_in_gdn=self.config.use_qk_norm_in_gdn, + compute_dtype=self.config.dtype, + ) + elif getattr(cfg, "use_gdn_kernel", False): + core_attn_out, next_recurrent_state = jax_chunk_gated_delta_rule( + query, + key, + value, + g, + beta, + chunk_size=cfg.gdn_chunk_size, + initial_state=recurrent_state, + use_qk_norm_in_gdn=cfg.use_qk_norm_in_gdn, + compute_dtype=cfg.dtype, + ) elif self.mesh is not None: logical_rules = get_logical_axis_rules() recurrent_state_arg = ( @@ -914,6 +1079,8 @@ def shard_mapped_delta_rule(q, k, v, g_val, beta_val, init_h): if model_mode != MODEL_MODE_TRAIN and active_cache is not None: active_cache.update_gdn_states(next_recurrent_state, next_conv_state) # pyrefly: ignore[bad-argument-type] + core_attn_out = checkpoint_name(core_attn_out, "context") + # ========================================================================= # STEP D: Final Output Stage # ========================================================================= @@ -1067,7 +1234,7 @@ def __init__(self, config: Config, mesh: Mesh, quant: None | Quant = None, *, rn cfg = self.config # 1. Instantiate and apply the routed experts block. - self.routed_experts = moe.RoutedMoE( + self.routed_experts = RoutedMoE( config=cfg, num_experts=cfg.num_experts, num_experts_per_tok=cfg.num_experts_per_tok, @@ -1137,86 +1304,220 @@ def __call__(self, hidden_states: Array, deterministic: bool) -> tuple[Array, Ar class Qwen3NextScannableBlock(nnx.Module): - """A scannable block of Qwen3-Next decoder layers. - - This module contains a fixed number of heterogeneous decoder layers that form - a repeating pattern, as defined by `config.inhomogeneous_layer_cycle_interval`. It is - intended to be the body of an `nn.scan` transformation to construct the full - decoder stack efficiently. + """A repeatable block of Qwen3-Next decoder layers, scanning local layers.""" - Attributes: - config: The model configuration object. - mesh: The device mesh for sharding. - model_mode: The operational mode (e.g., 'train', 'prefill'). - quant: Optional quantization configuration. - """ + def __init__( + self, + config: Config, + mesh: Mesh, + model_mode: str, + rngs: nnx.Rngs, + quant: None | Quant = None, + num_of_layers: int | None = None, + remat_policy_fn: Any = None, + apply_internal_remat: bool = False, + ): + """Initializes the instance. - def __init__(self, config: Config, mesh: Mesh, model_mode: str, quant: None | Quant = None, *, rngs: nnx.Rngs): + Args: + config: The Config object with model hyperparameters. + mesh: The device mesh for distributed training. + model_mode: One of MODEL_MODE_TRAIN, MODEL_MODE_PREFILL, or MODEL_MODE_AUTOREGRESSIVE. + rngs: The random number generators for initialization. + quant: The quantization configuration. + num_of_layers: The number of layers in the block. + remat_policy_fn: The resolved rematerialization policy function. + apply_internal_remat: When True, the block rematerializes its own local + (scanned) and global layers, and the caller must NOT also apply + block-level remat. + """ self.config = config self.mesh = mesh self.model_mode = model_mode self.quant = quant self.rngs = rngs - cfg = self.config + cycle_interval = config.inhomogeneous_layer_cycle_interval + if num_of_layers is None: + num_of_layers = cycle_interval + self.num_of_layers = num_of_layers + self.remat_policy_fn = remat_policy_fn + self.apply_internal_remat = apply_internal_remat + + if not 0 <= num_of_layers <= cycle_interval: + raise ValueError( + f"Qwen3NextScannableBlock must contain between 0 and {cycle_interval} layers; got {num_of_layers}." + ) + + # Calculate local (GatedDeltaNet) vs global (FullAttention) layer counts for the block. + self.num_local = sum(1 for i in range(num_of_layers) if (i + 1) % cycle_interval != 0) + self.num_global = sum(1 for i in range(num_of_layers) if (i + 1) % cycle_interval == 0) + + if self.num_local > 0: + self.local_layers = nnx_scan.create_scanned_layers( + lambda layer_rngs: Qwen3NextDecoderLayer( + config=self.config, + mesh=self.mesh, + model_mode=self.model_mode, + quant=self.quant, + layer_idx=0, # layer_idx 0 is a GatedDeltaNet layer + rngs=layer_rngs, + ), + length=self.num_local, + param_scan_axis=self.config.param_scan_axis, + metadata_axis_name="local_layers", + rngs=self.rngs, + ) + else: + self.local_layers = None - # Instantiate each layer within the block in __init__ - for i in range(cfg.inhomogeneous_layer_cycle_interval): - layer_rngs = self.rngs.fork() # Fork RNGs for each layer - layer_name = f"layer_{i}" - layer = Qwen3NextDecoderLayer( + if self.num_global > 0: + self.global_layer = Qwen3NextDecoderLayer( config=self.config, mesh=self.mesh, - quant=self.quant, model_mode=self.model_mode, - layer_idx=i, - rngs=layer_rngs, + quant=self.quant, + layer_idx=cycle_interval - 1, # layer_idx cycle_interval-1 is a FullAttention layer + rngs=self.rngs, ) - setattr(self, layer_name, layer) + else: + self.global_layer = None + + def _run_layer(self, layer, y, layer_kwargs, kv_cache=None): + """Invokes one ``Qwen3NextDecoderLayer``, returning ``(output, updated_kv_cache)``.""" + out = layer(y, **layer_kwargs, kv_cache=kv_cache) + return out if isinstance(out, tuple) else (out, None) + + @property + def _remat_enabled(self): + """Whether the block rematerializes its own layers.""" + return self.apply_internal_remat and self.config.remat_policy != "none" + + def _scan_local_layers(self, y, layer_kwargs): + """Runs the local (linear attention / GatedDeltaNet) layers via a per-layer rematerialized ``jax.lax.scan``.""" + remat = self._remat_enabled + return nnx_scan.apply_scanned_layers( + self.local_layers, + y, + length=self.num_local, + param_scan_axis=self.config.param_scan_axis, + apply_fn=lambda layer, carry: self._run_layer(layer, carry, layer_kwargs)[0], + remat=remat, + remat_policy=self.remat_policy_fn if remat else None, + prevent_cse=maxtext_utils.should_prevent_cse_in_remat(self.config) if remat else True, + ) + + def _scan_global_layer(self, y, layer_kwargs): + """Runs the single global-attention layer inside a length-1 ``jax.lax.scan``.""" + cfg = self.config + graphdef_g, intermediate_g, other_g = nnx.split(self.global_layer, nnx.Intermediate, ...) + intermediate_xs = jax.tree.map(lambda x: x[None], intermediate_g) + + def run_global_layer(carry, intermediate_slice): + hidden_states, other = carry + layer = nnx.merge(graphdef_g, intermediate_slice, other) + new_hidden_states = self._run_layer(layer, hidden_states, layer_kwargs)[0] + _, new_intermediate, new_other = nnx.split(layer, nnx.Intermediate, ...) + return (new_hidden_states, new_other), new_intermediate + + global_remat_policy = self.remat_policy_fn + offload_names = maxtext_utils.get_save_and_offload_names(cfg) + if offload_names[0] or offload_names[1]: + save_names, offload_to_device = offload_names + global_remat_policy = jax.checkpoint_policies.save_only_these_names(*(save_names + offload_to_device)) + + if self._remat_enabled: + prevent_cse = maxtext_utils.should_prevent_cse_in_remat(self.config) + run_global_layer = jax.checkpoint( + run_global_layer, + policy=global_remat_policy, + prevent_cse=prevent_cse, + ) + + with xla_metadata.set_xla_metadata(**{"skip-simplify-while-loops_trip-count-one": "true"}): + (y, final_other), stacked_intermediate = jax.lax.scan( + run_global_layer, + (y, other_g), + intermediate_xs, + length=1, + ) + + intermediate_state = jax.tree.map(lambda x: x[0], stacked_intermediate) + nnx.update(self.global_layer, final_other, intermediate_state) + return y + + def _forward_with_external_kv_cache(self, y, kv_cache, layer_kwargs): + """Runs the block with externally-supplied per-layer kv caches (vLLM PagedAttention / Mamba).""" + updated_kvs = [] + + if self.local_layers is not None: + graphdef, params, state = nnx.split(self.local_layers, nnx.Param, ...) + scan_axis = self.config.param_scan_axis + if scan_axis != 0: + params = jax.tree.map(lambda x: jnp.moveaxis(x, scan_axis, 0), params) + per_layer_states = [] + for i in range(self.num_local): + current_params = jax.tree.map(lambda x, i=i: x[i], params) + current_state = jax.tree.map(lambda x, i=i: x[i], state) + layer = nnx.merge(graphdef, current_params, current_state) + current_kv = kv_cache[i] if (kv_cache is not None and i < len(kv_cache)) else None + y, new_kv = self._run_layer(layer, y, layer_kwargs, current_kv) + updated_kvs.append(new_kv) + per_layer_states.append(nnx.state(layer)) + + stacked_state = jax.tree.map(lambda *xs: jnp.stack(xs), *per_layer_states) + if scan_axis != 0: + stacked_params, stacked_other = stacked_state.split(nnx.Param, ...) + stacked_params = jax.tree.map(lambda x: jnp.moveaxis(x, 0, scan_axis), stacked_params) + stacked_state = nnx.State.merge(stacked_params, stacked_other) + nnx.update(self.local_layers, stacked_state) + + if self.global_layer is not None: + global_kv = kv_cache[self.num_local] if (kv_cache is not None and self.num_local < len(kv_cache)) else None + y, new_kv = self._run_layer(self.global_layer, y, layer_kwargs, global_kv) + updated_kvs.append(new_kv) + + return y, tuple(updated_kvs) def __call__( self, - carry: jnp.ndarray, + inputs: jnp.ndarray, decoder_segment_ids: None | jnp.ndarray, decoder_positions: None | jnp.ndarray, deterministic: bool, model_mode: str, previous_chunk=None, slot: None | int = None, + page_state=None, + bidirectional_mask=None, kv_cache=None, attention_metadata=None, ) -> tuple[Array, None]: - """Applies the block of decoder layers to the input carry. + cfg = self.config + inputs = nn.with_logical_constraint(inputs, ("activation_batch", "activation_norm_length", "activation_embed")) + inputs = checkpoint_name(inputs, "decoder_layer_input") - Args: - carry: The input tensor from the previous scan iteration. - # ... other arguments are broadcasted to each iteration. + layer_kwargs = { + "decoder_segment_ids": decoder_segment_ids, + "decoder_positions": decoder_positions, + "deterministic": deterministic, + "model_mode": model_mode, + "slot": slot, + "previous_chunk": previous_chunk, + "attention_metadata": attention_metadata, + } - Returns: - A tuple containing the output of the block (the new carry) and an empty - value for the scan's `y` collection. - """ - cfg = self.config - x = carry - - # Loop over the number of sub-layers that make up one repeating pattern. - for i in range(cfg.inhomogeneous_layer_cycle_interval): - layer = getattr(self, f"layer_{i}") - # The second return value is kv_cache, which we ignore here because - # it is not passed as a carry in scannable layers. - x, _ = layer( - x, - decoder_segment_ids, - decoder_positions, - deterministic, - model_mode, - previous_chunk, - slot, - kv_cache=kv_cache, - attention_metadata=attention_metadata, - ) + if kv_cache is not None: + return self._forward_with_external_kv_cache(inputs, kv_cache, layer_kwargs) + + y = inputs + if self.local_layers is not None: + y = self._scan_local_layers(y, layer_kwargs) + if self.global_layer is not None: + y = self._scan_global_layer(y, layer_kwargs) - # The output of the block is the carry for the next scan iteration. - return x, None + if cfg.scan_layers: + return y, None + return y class Qwen3NextDecoderLayer(nnx.Module): @@ -1289,6 +1590,15 @@ def __init__( # Instantiate our `Qwen3NextSparseMoeBlock`. self.mlp = Qwen3NextSparseMoeBlock(config=cfg, mesh=self.mesh, quant=self.quant, rngs=rngs) + self.is_mhc_enabled = getattr(cfg, "mhc_expansion_rate", 1) > 1 + if self.is_mhc_enabled: + self.mhc_attention = mhc.ManifoldConstrainedHyperConnections( + config=cfg, dim=cfg.emb_dim, mesh=self.mesh, rngs=rngs + ) + self.mhc_mlp = mhc.ManifoldConstrainedHyperConnections( + config=cfg, dim=cfg.emb_dim, mesh=self.mesh, rngs=rngs + ) + def __call__( self, inputs: jnp.ndarray, @@ -1304,6 +1614,61 @@ def __call__( # Unpack inputs if it's a tuple (e.g. from a previous layer returning (hidden_states, kv_cache)) if isinstance(inputs, tuple): inputs = inputs[0] + + if self.is_mhc_enabled: + mhc_expand, mhc_reduce = mhc.get_functions(self.config.mhc_expansion_rate) + inputs = mhc_expand(inputs) + new_kv_cache = None + + def attention_branch(inputs): + nonlocal new_kv_cache + if isinstance(self.attention, Qwen3NextFullAttention): + out, new_kv_cache = cast(Qwen3NextFullAttention, self.attention)( + inputs, + decoder_segment_ids, + decoder_positions, + deterministic, + model_mode, + kv_cache=kv_cache, + attention_metadata=attention_metadata, + ) + else: + out, new_kv_cache = cast(Qwen3NextGatedDeltaNet, self.attention)( + inputs, + model_mode=model_mode, + kv_cache=kv_cache, + decoder_segment_ids=decoder_segment_ids, + attention_metadata=attention_metadata, + ) + return out + + intermediate_inputs, _ = self.mhc_attention( + self.input_layernorm, + attention_branch, + x=inputs, + mhc_type=HyperConnectionType.MLP_DENSE, + ) + + def mlp_branch(inputs): + mlp_output, load_balance_loss = self.mlp(inputs, deterministic=deterministic) + if self.config.load_balance_loss_weight > 0.0 and load_balance_loss is not None: + self.moe_lb_loss = nnx.Intermediate(load_balance_loss) + return mlp_output + + layer_output, _ = self.mhc_mlp( + self.post_attention_layernorm, + mlp_branch, + x=intermediate_inputs, + mhc_type=HyperConnectionType.MLP_DENSE, + ) + + layer_output = mhc_reduce(layer_output) + layer_output = nn.with_logical_constraint( + layer_output, + self.activation_axis_names, + ) + return layer_output, new_kv_cache + residual = inputs # First LayerNorm, applied before the attention block. @@ -2131,17 +2496,22 @@ def __call__( x, _ = self.patch_embed(hidden_states) x = x.reshape(batch_size, -1, self.config.hidden_size_for_vit) + valid_grid = None if attention_mask is not None and video_grid_thw is None: raise ValueError("video_grid_thw is required when video_mask is provided.") - pos = self.pos_embed_interpolate( - num_frames, - height, - width, - video_grid_thw=video_grid_thw, # pyrefly: ignore[bad-argument-type] - attention_mask=attention_mask, - ) + if attention_mask is not None and batch_size != 1: + raise ValueError("Padded Qwen3-Omni vision encoding currently supports batch size one.") + if video_grid_thw is not None: + grid = video_grid_thw[0] if getattr(video_grid_thw, "ndim", 1) == 2 else video_grid_thw + valid_grid = tuple(int(dim) for dim in grid) + pos = self.pos_embed_interpolate(num_frames, height, width) + if attention_mask is not None and valid_grid is not None: + valid_pos = self.pos_embed_interpolate(*valid_grid) + valid_indices = jnp.nonzero(attention_mask[0], size=math.prod(valid_grid))[0] + pos = jnp.zeros_like(pos).at[valid_indices].set(valid_pos) + + pos = pos[jnp.newaxis, :, :] x = x + pos - valid_grid = video_grid_thw h_traj = [] for i in range(self.depth): diff --git a/src/maxtext/optimizers/optimizers.py b/src/maxtext/optimizers/optimizers.py index 67e1f589ca..de288b24a2 100644 --- a/src/maxtext/optimizers/optimizers.py +++ b/src/maxtext/optimizers/optimizers.py @@ -204,7 +204,7 @@ def get_optimizer(config, learning_rate_schedule, model=None): ns_steps = 10 else: ns_coeffs = (3.4445, -4.7750, 2.0315) - ns_steps = 5 + ns_steps = getattr(config, "muon_ns_steps", 5) muon_kwargs = { # Shared parameters: "nesterov" uses default diff --git a/src/maxtext/utils/muon_utils.py b/src/maxtext/utils/muon_utils.py index ff77c57807..eb2fe18d6c 100644 --- a/src/maxtext/utils/muon_utils.py +++ b/src/maxtext/utils/muon_utils.py @@ -101,7 +101,7 @@ def transform_logic(path: Tuple[str, ...]) -> Optional[mdn]: # 2 Special weights # 2.1 Special weights: MoE, [0, L, -2, -1] # L (optional) stands for layer when scan_layers=True - if "MoeBlock_0" in path: + if "MoeBlock_0" in path or "routed_experts" in path: # exclude gate if _is_path_contain_any(("wi_0", "wi_1", "wo"), path): return mdn((-2,), (-1,)) diff --git a/tests/unit/nnx_decoders_test.py b/tests/unit/nnx_decoders_test.py index fcd5acb5cc..50f47aa8ff 100644 --- a/tests/unit/nnx_decoders_test.py +++ b/tests/unit/nnx_decoders_test.py @@ -51,7 +51,7 @@ from maxtext.layers.embeddings import Embed from maxtext.layers.nnx_decoders import NNXDecoder, NNXDecoderLayer, deepstack_process from maxtext.layers.normalizations import RMSNorm -from maxtext.models import gemma4, gemma4_small +from maxtext.models import gemma4, gemma4_small, qwen3 from maxtext.models.gpt3 import Gpt3LayerNorm from maxtext.models.llama2 import LlamaDecoderLayer from maxtext.utils import maxtext_utils @@ -716,8 +716,219 @@ def test_scan_layers(self): self.assertEqual(logits.shape, (batch, seq_len, cfg.vocab_size)) -if __name__ == "__main__": - unittest.main() +class _StatefulGemma4DecoderLayer(nnx.Module): + """Small stand-in that exposes cache ordering and mutable-state updates.""" + + def __init__(self, *, attention_type, **unused_kwargs): + self.increment = 10 if attention_type == AttentionType.GLOBAL else 1 + self.call_count = nnx.Intermediate(jnp.array(0, dtype=jnp.int32)) + self.received_attention_metadata = nnx.Intermediate(jnp.array(False)) + + def __call__( + self, + inputs, + *unused_args, + kv_cache=None, + attention_metadata=None, + **unused_kwargs, + ): + self.call_count.value += 1 + self.received_attention_metadata.value = attention_metadata is not None + output = inputs + self.increment + if kv_cache is None: + return output + return output, kv_cache + self.increment + + +class _SowingGemma4DecoderLayer(nnx.Module): + """Stand-in whose global layer sows an accumulating Intermediate, like MoE moe_lb_loss.""" + + def __init__(self, *, attention_type, **unused_kwargs): + self.is_global = attention_type == AttentionType.GLOBAL + # A trivial variable so the local layers have state for apply_scanned_layers to + # scan over (a bare module has nothing to scan and lax.scan can't infer length). + self.marker = nnx.Intermediate(jnp.zeros(())) + + def __call__(self, inputs, *unused_args, kv_cache=None, **unused_kwargs): + output = inputs + 1 + if self.is_global: + # nnx.sow appends into a tuple by default, so it grows across calls -- the + # MoE moe_lb_loss pattern that must not enter the global length-1 scan carry. + self.sow(nnx.Intermediate, "moe_lb_loss", jnp.sum(output)) + if kv_cache is None: + return output + return output, kv_cache + + +class TestGemma4ScannableBlock(unittest.TestCase): + """Tests Gemma4's nested local/global decoder block behavior.""" + + def setUp(self): + super().setUp() + self.config = SimpleNamespace( + dtype=jnp.float32, + param_scan_axis=1, + remat_policy="none", + scan_layers=True, + ) + + def _make_block(self): + return gemma4.Gemma4ScannableBlock( + config=self.config, + mesh=None, + model_mode=MODEL_MODE_AUTOREGRESSIVE, + rngs=nnx.Rngs(0), + ) + + def test_updates_state_through_global_single_iteration_scan(self): + with mock.patch.object(gemma4, "Gemma4DecoderLayer", _StatefulGemma4DecoderLayer): + block = self._make_block() + output, updated_kvs = block( + jnp.zeros((1, 1, 1)), + decoder_segment_ids=None, + decoder_positions=None, + deterministic=True, + model_mode=MODEL_MODE_AUTOREGRESSIVE, + ) + + np.testing.assert_array_equal(output, jnp.full((1, 1, 1), 15)) + self.assertIsNone(updated_kvs) + np.testing.assert_array_equal(block.local_layers.call_count.value, jnp.ones(5, dtype=jnp.int32)) + np.testing.assert_array_equal(block.global_layer.call_count.value, 1) + + def test_global_layer_sown_intermediate_accumulates_across_calls(self): + """A global layer that sows an accumulating Intermediate (e.g. MoE moe_lb_loss) + must not break the length-1 scan carry, even when the Intermediate already + exists from a previous call and the sow grows its tuple (1 -> 2 elements).""" + call_kwargs = { + "decoder_segment_ids": None, + "decoder_positions": None, + "deterministic": True, + "model_mode": MODEL_MODE_AUTOREGRESSIVE, + } + with mock.patch.object(gemma4, "Gemma4DecoderLayer", _SowingGemma4DecoderLayer): + block = self._make_block() + # First call creates moe_lb_loss on the global layer (1-tuple). + block(jnp.zeros((1, 1, 1)), **call_kwargs) + # Second call: moe_lb_loss already exists and the sow appends -> 2-tuple. + # Carrying it in the scan would change the carry pytree; the type-based + # split keeps Intermediates on the ys path instead. + block(jnp.zeros((1, 1, 1)), **call_kwargs) + + self.assertEqual(len(block.global_layer.moe_lb_loss.value), 2) + + def test_restores_local_state_and_preserves_kv_order(self): + attention_metadata = object() + + with mock.patch.object(gemma4, "Gemma4DecoderLayer", _StatefulGemma4DecoderLayer): + block = self._make_block() + output, updated_kvs = block( + jnp.zeros((1, 1, 1)), + decoder_segment_ids=None, + decoder_positions=None, + deterministic=True, + model_mode=MODEL_MODE_AUTOREGRESSIVE, + kv_cache=tuple(jnp.array(i) for i in range(6)), + attention_metadata=attention_metadata, + ) + + np.testing.assert_array_equal(output, jnp.full((1, 1, 1), 15)) + np.testing.assert_array_equal(jnp.stack(updated_kvs), jnp.array([1, 2, 3, 4, 5, 15])) + np.testing.assert_array_equal(block.local_layers.call_count.value, jnp.ones(5, dtype=jnp.int32)) + np.testing.assert_array_equal( + block.local_layers.received_attention_metadata.value, + jnp.ones(5, dtype=jnp.bool_), + ) + np.testing.assert_array_equal(block.global_layer.call_count.value, 1) + np.testing.assert_array_equal(block.global_layer.received_attention_metadata.value, True) + + +class _StatefulQwen3NextDecoderLayer(nnx.Module): + """Small stand-in that exposes cache ordering and mutable-state updates for Qwen3-Next.""" + + def __init__(self, *, layer_idx, **unused_kwargs): + is_global = (layer_idx + 1) % 4 == 0 + self.increment = 10 if is_global else 1 + self.call_count = nnx.Intermediate(jnp.array(0, dtype=jnp.int32)) + self.received_attention_metadata = nnx.Intermediate(jnp.array(False)) + + def __call__( + self, + inputs, + *unused_args, + kv_cache=None, + attention_metadata=None, + **unused_kwargs, + ): + self.call_count.value += 1 + self.received_attention_metadata.value = attention_metadata is not None + output = inputs + self.increment + if kv_cache is None: + return output + return output, kv_cache + self.increment + + +class TestQwen3NextScannableBlock(unittest.TestCase): + """Tests Qwen3-Next's nested local/global decoder block behavior.""" + + def setUp(self): + super().setUp() + self.config = SimpleNamespace( + dtype=jnp.float32, + param_scan_axis=1, + remat_policy="none", + scan_layers=True, + inhomogeneous_layer_cycle_interval=4, + ) + + def _make_block(self): + return qwen3.Qwen3NextScannableBlock( + config=self.config, + mesh=None, + model_mode=MODEL_MODE_AUTOREGRESSIVE, + rngs=nnx.Rngs(0), + ) + + def test_updates_state_through_global_single_iteration_scan(self): + with mock.patch.object(qwen3, "Qwen3NextDecoderLayer", _StatefulQwen3NextDecoderLayer): + block = self._make_block() + output, updated_kvs = block( + jnp.zeros((1, 1, 1)), + decoder_segment_ids=None, + decoder_positions=None, + deterministic=True, + model_mode=MODEL_MODE_AUTOREGRESSIVE, + ) + + np.testing.assert_array_equal(output, jnp.full((1, 1, 1), 13)) + self.assertIsNone(updated_kvs) + np.testing.assert_array_equal(block.local_layers.call_count.value, jnp.ones(3, dtype=jnp.int32)) + np.testing.assert_array_equal(block.global_layer.call_count.value, 1) + + def test_restores_local_state_and_preserves_kv_order(self): + attention_metadata = object() + + with mock.patch.object(qwen3, "Qwen3NextDecoderLayer", _StatefulQwen3NextDecoderLayer): + block = self._make_block() + output, updated_kvs = block( + jnp.zeros((1, 1, 1)), + decoder_segment_ids=None, + decoder_positions=None, + deterministic=True, + model_mode=MODEL_MODE_AUTOREGRESSIVE, + kv_cache=tuple(jnp.array(i) for i in range(4)), + attention_metadata=attention_metadata, + ) + + np.testing.assert_array_equal(output, jnp.full((1, 1, 1), 13)) + np.testing.assert_array_equal(jnp.stack(updated_kvs), jnp.array([1, 2, 3, 13])) + np.testing.assert_array_equal(block.local_layers.call_count.value, jnp.ones(3, dtype=jnp.int32)) + np.testing.assert_array_equal( + block.local_layers.received_attention_metadata.value, + jnp.ones(3, dtype=jnp.bool_), + ) + np.testing.assert_array_equal(block.global_layer.call_count.value, 1) + np.testing.assert_array_equal(block.global_layer.received_attention_metadata.value, True) class _StatefulGemma4DecoderLayer(nnx.Module): @@ -1259,3 +1470,7 @@ def mock_donor_idx(lyr, layer_types, num_kv_shared): model_mode=MODEL_MODE_TRAIN, kv_caches=kv_caches, ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/unit/param_mapping_test.py b/tests/unit/param_mapping_test.py index cea6485817..f8fd22c5d1 100644 --- a/tests/unit/param_mapping_test.py +++ b/tests/unit/param_mapping_test.py @@ -109,7 +109,14 @@ def test_qwen3_next_mapping_scanned(self): maxtext_config = mock.Mock() maxtext_config.inhomogeneous_layer_cycle_interval = 2 mapping = param_mapping.QWEN3_NEXT_MAXTEXT_TO_HF_PARAM_MAPPING(config, maxtext_config, scan_layers=True) - self.assertIn("params-decoder-layers-layer_0-input_layernorm-scale", mapping) + self.assertIn("params-decoder-scanned_blocks-local_layers-input_layernorm-scale", mapping) + self.assertIn("params-decoder-scanned_blocks-global_layer-input_layernorm-scale", mapping) + num_blocks = config["num_hidden_layers"] // maxtext_config.inhomogeneous_layer_cycle_interval + local_val = mapping["params-decoder-scanned_blocks-local_layers-input_layernorm-scale"] + global_val = mapping["params-decoder-scanned_blocks-global_layer-input_layernorm-scale"] + self.assertEqual(len(local_val), num_blocks) + self.assertEqual(len(local_val[0]), 1) + self.assertEqual(len(global_val), num_blocks) def test_deepseek_mapping(self): config = { diff --git a/tests/unit/qwen3_next_vs_reference_test.py b/tests/unit/qwen3_next_vs_reference_test.py index e9efad376f..7ffa0c0990 100644 --- a/tests/unit/qwen3_next_vs_reference_test.py +++ b/tests/unit/qwen3_next_vs_reference_test.py @@ -22,6 +22,7 @@ import jax import jax.numpy as jnp from jax.sharding import Mesh +from jax.test_util import check_grads from maxtext.configs import pyconfig from maxtext.layers import normalizations from maxtext.layers.normalizations import Qwen3NextRMSNorm, Qwen3NextRMSNormGated @@ -1037,6 +1038,40 @@ def run_jax(x): ) print("test_qwen3_next_sparse_moe_block passed!") + def test_invert_unit_lower_triangular_log_depth(self): + """Test for loss at chunk_size 256.""" + jax.config.update("jax_enable_x64", True) # Use float64 for precise testing + chunk_size = 256 + + # Generate a random matrix and make it strictly lower triangular + key = jax.random.PRNGKey(chunk_size) + S_random = jax.random.normal(key, (chunk_size, chunk_size), dtype=jnp.float64) / chunk_size + S = jnp.tril(S_random, k=-1) + + # The matrix to invert is (I + S) + identity = jnp.eye(chunk_size, dtype=jnp.float64) + matrix_to_invert = identity + S + + # Using our custom function + A = qwen3.invert_unit_lower_triangular_log_depth(S) + + # The product A @ (I + S) should be exactly the identity matrix + # Wait, due to numerical precision, we should check for max error (loss) + reconstructed_identity = A @ matrix_to_invert + + # Compute loss for forward pass + loss = jnp.max(jnp.abs(reconstructed_identity - identity)) + + # We expect the loss to be very small, around numerical precision + self.assertLess(loss, 1e-10, f"Failed for chunk_size {chunk_size} with loss {loss}") + + # Verify backward pass accuracy using jax.test_util.check_grads + # This uses finite differences to check the correctness of the custom VJP + # We check the gradients for the function. + # `check_grads` will assert if finite difference gradients + # don't match the custom VJP gradients. + check_grads(qwen3.invert_unit_lower_triangular_log_depth, (S,), order=1, modes=["rev"]) + def test_gated_delta_net_full(self): """Tests the full Qwen3NextGatedDeltaNet layer for numerical correctness.""" print("Running test_gated_delta_net_full...")