diff --git a/docs/reference/core_concepts/moe_configuration.md b/docs/reference/core_concepts/moe_configuration.md index 4f6945d425..faa5706aa7 100644 --- a/docs/reference/core_concepts/moe_configuration.md +++ b/docs/reference/core_concepts/moe_configuration.md @@ -93,6 +93,8 @@ MaxText implements an exact, paper-aligned version of DeepSeek V4's load balanci `use_gmm_v2`: If enabled, use the Tokamax GMM v2 kernel for grouped matrix multiplication. Requires `use_tokamax_gmm` to be True. +`use_gmm_v2_tiling_fn`: If enabled, use the heuristic tiling function from Tokamax GMM v2. Recommended when not using custom tile sizes. + `megablox`: If enabled, use Megablox for sparse matrix operations. Effective only when `use_tokamax_gmm` is False. `capacity_factor`: A scalar multiplier for expert capacity. Effective only when `sparse_matmul` is False. @@ -155,5 +157,5 @@ Implementation Support: - Tokamax Ragged Dot (Includes two implementations): - **GMM v1**: Uses Tokamax's native autotuner; does not accept manual tile sizes from MaxText. - - **GMM v2**: Supports all 18 manual tiling configurations. + - **GMM v2**: Supports all 18 manual tiling configurations. Optionally, use `use_gmm_v2_tiling_fn=True` for heuristic tiling. - Enabled for FP8 and BF16. diff --git a/src/maxtext/configs/base.yml b/src/maxtext/configs/base.yml index a4cf6877e3..974772f59c 100644 --- a/src/maxtext/configs/base.yml +++ b/src/maxtext/configs/base.yml @@ -274,6 +274,9 @@ merge_gating_gmm: false use_tokamax_gmm: false # Whether to use Tokamax GMM v2 for MoE kernel. Requires use_tokamax_gmm=true. use_gmm_v2: false +# Whether to use the heuristic tiling function from Tokamax GMM v2, when use_gmm_v2=true. +# Recommended when not using custom tile sizes. +use_gmm_v2_tiling_fn: false norm_topk_prob: false # boolean to enable the top-k probability normalization. qwen3-specific normalization of router weights. diff --git a/src/maxtext/configs/types.py b/src/maxtext/configs/types.py index 03e5db969b..084679455c 100644 --- a/src/maxtext/configs/types.py +++ b/src/maxtext/configs/types.py @@ -1037,6 +1037,11 @@ class MoEKernels(BaseModel): description="Whether to use Tokamax GMM v2 for MoE kernel.", ) + use_gmm_v2_tiling_fn: bool = Field( + False, + description="Whether to use the tiling function from Tokamax GMM v2, when use_gmm_v2=true.", + ) + class DeepSeekMoE(BaseModel): """Configuration specific to DeepSeek-style MoE layers.""" @@ -4183,6 +4188,11 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de if self.use_gmm_v2: if not self.use_tokamax_gmm: raise ValueError("GMM v2 requires `use_tokamax_gmm=True`.") + if not self.use_gmm_v2_tiling_fn: + logger.info( + "A heuristic tiling function is available for Tokamax GMM v2 and can be enabled with `use_gmm_v2_tiling_fn=True`. " + "This is recommended when not using custom tile sizes." + ) if self.use_batch_split_schedule: raise ValueError("GMM v2 is not supported with a batch split schedule.") diff --git a/src/maxtext/kernels/megablox/ops.py b/src/maxtext/kernels/megablox/ops.py index c717eda455..ffc5b37c4a 100644 --- a/src/maxtext/kernels/megablox/ops.py +++ b/src/maxtext/kernels/megablox/ops.py @@ -74,6 +74,7 @@ def gmm( qwix_rule: qwix.QtRule | None = None, use_manual_quantization: bool = False, # used in batchsplit use_gmm_v2: bool = False, + use_gmm_v2_tiling_fn: bool = False, partial_sum: jnp.ndarray | None = None, ): """Grouped matrix multiplication operation.""" @@ -105,7 +106,7 @@ def gmm( gmm_fwd_bwd = lambda *args: _gmm_fwd(*args)[0] # pylint: disable=C3001 gmm_fwd_bwd = jax.custom_vjp( gmm_fwd_bwd, - nondiff_argnums=(3, 4, 7, 8, 9, 10, 11, 12, 13, 14, 15), + nondiff_argnums=(3, 4, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16), ) gmm_fwd_bwd.defvjp(_gmm_fwd, functools.partial(_gmm_bwd, lhs.dtype, rhs.dtype)) return gmm_fwd_bwd( @@ -125,6 +126,7 @@ def gmm( lhs_vma_axes, rhs_vma_axes, use_gmm_v2, + use_gmm_v2_tiling_fn, partial_sum, ) @@ -161,6 +163,7 @@ def _gmm_fwd( lhs_vma_axes: tuple = tuple(), rhs_vma_axes: tuple = tuple(), use_gmm_v2: bool = False, + use_gmm_v2_tiling_fn: bool = False, partial_sum: jnp.ndarray | None = None, ) -> tuple[ jnp.ndarray, @@ -207,6 +210,7 @@ def _gmm_fwd( group_sizes, preferred_element_type, tiling, + use_gmm_v2_tiling_fn, group_offset, partial_sum, transpose_rhs, @@ -366,6 +370,7 @@ def _fwd_run_tokamax_v2( group_sizes: jnp.ndarray, preferred_element_type: jnp.dtype, tiling: tuple, + use_gmm_v2_tiling_fn: bool, group_offset: jnp.ndarray | None, partial_sum: jnp.ndarray | None, transpose_rhs: bool, @@ -382,18 +387,18 @@ def _fwd_run_tokamax_v2( rhs_operand = rhs_operand.qvalue rhs_scale = _fwd_prepare_rhs_scale(rhs, transpose_rhs=transpose_rhs) - custom_fwd_tiling = gmm_v2.TileSizes( - tile_m=tiling[0], - tile_k=tiling[1], - tile_n=tiling[2], - ) + # Heuristic to calculate tiling + fwd_tiling = gmm_v2.calculate_tiling + if not use_gmm_v2_tiling_fn: + # Custom tiling + fwd_tiling = gmm_v2.TileSizes(tile_m=tiling[0], tile_k=tiling[1], tile_n=tiling[2]) return gmm_v2.gmm_v2( lhs=lhs, # pyrefly: ignore[bad-argument-type] rhs=rhs_operand, # pyrefly: ignore[bad-argument-type] group_sizes=group_sizes, rhs_scale=rhs_scale, - tile_info=custom_fwd_tiling, + tile_info=fwd_tiling, preferred_element_type=preferred_element_type, partial_sum=partial_sum, group_offset=group_offset, @@ -449,6 +454,7 @@ def _gmm_bwd( lhs_vma_axes: tuple, rhs_vma_axes: tuple, use_gmm_v2: bool, + use_gmm_v2_tiling_fn: bool, residual: tuple[ jnp.ndarray | qpl.QArray, jnp.ndarray | qpl.QArray, @@ -495,6 +501,7 @@ def _gmm_bwd( use_manual_quantization, interpret, lhs_vma_axes, + use_gmm_v2_tiling_fn, ) # 4. DRHS Gradient Execution @@ -513,6 +520,7 @@ def _gmm_bwd( interpret, rhs_vma_axes, quantization_rule, + use_gmm_v2_tiling_fn, ) # 5. Output Formatting @@ -622,6 +630,7 @@ def _compute_dlhs( use_manual_quantization: bool, interpret: bool, lhs_vma_axes: tuple, + use_gmm_v2_tiling_fn: bool, ) -> jnp.ndarray: """Routes execution of DLHS based on backend choices.""" if use_tokamax_backend and not use_gmm_v2: @@ -634,7 +643,9 @@ def _compute_dlhs( use_manual_quantization, ) elif use_tokamax_backend and use_gmm_v2: - return _dlhs_run_tokamax_v2(dlhs_dout, rhs, group_sizes, group_offset, lhs_dtype, tiling, transpose_rhs) + return _dlhs_run_tokamax_v2( + dlhs_dout, rhs, group_sizes, group_offset, lhs_dtype, tiling, use_gmm_v2_tiling_fn, transpose_rhs + ) else: return _dlhs_run_megablox( dlhs_dout, rhs, group_sizes, group_offset, lhs_dtype, tiling, transpose_rhs, interpret, lhs_vma_axes @@ -707,6 +718,7 @@ def _dlhs_run_tokamax_v2( group_offset: jnp.ndarray | None, lhs_dtype: jax.typing.DTypeLike, tiling: tuple, + use_gmm_v2_tiling_fn: bool, transpose_rhs: bool, ) -> jnp.ndarray: """Executes Tokamax GMM V2 backend for DLHS = DLHS_dout @ RHS^T.""" @@ -714,7 +726,11 @@ def _dlhs_run_tokamax_v2( dlhs_rhs = rhs if transpose_rhs else rhs.swapaxes(1, 2) dlhs_lhs = dlhs_dout.qvalue if isinstance(dlhs_dout, qpl.QArray) else dlhs_dout - custom_dlhs_tiling = gmm_v2.TileSizes(tile_m=tiling[3], tile_k=tiling[4], tile_n=tiling[5]) + # Heuristic to calculate GMM V2 tiling + dlhs_tiling = gmm_v2.calculate_tiling + if not use_gmm_v2_tiling_fn: + # Custom tiling for DLHS + dlhs_tiling = gmm_v2.TileSizes(tile_m=tiling[3], tile_k=tiling[4], tile_n=tiling[5]) dlhs = gmm_v2.gmm_v2( lhs=dlhs_lhs, @@ -722,7 +738,7 @@ def _dlhs_run_tokamax_v2( group_sizes=group_sizes, # rhs scale is already applied to dlhs_lhs rhs_scale=None, - tile_info=custom_dlhs_tiling, + tile_info=dlhs_tiling, preferred_element_type=lhs_dtype, # pyrefly: ignore[bad-argument-type] group_offset=group_offset, ) @@ -778,12 +794,15 @@ def _compute_drhs( interpret: bool, rhs_vma_axes: tuple, quantization_rule: qwix.QtRule | None, + use_gmm_v2_tiling_fn: bool, ) -> jnp.ndarray: """Routes execution of DRHS based on backend choices.""" if use_tokamax_backend and not use_gmm_v2: drhs = _drhs_run_tokamax_v1(drhs_dout, lhs, group_sizes, rhs_dtype, use_manual_quantization) elif use_tokamax_backend and use_gmm_v2: - drhs = _drhs_run_tokamax_v2(drhs_dout, lhs, group_sizes, group_offset, num_actual_groups, rhs_dtype, tiling) + drhs = _drhs_run_tokamax_v2( + drhs_dout, lhs, group_sizes, group_offset, num_actual_groups, rhs_dtype, tiling, use_gmm_v2_tiling_fn + ) else: drhs = _drhs_run_megablox( drhs_dout, lhs, group_sizes, group_offset, num_actual_groups, rhs_dtype, tiling, interpret, rhs_vma_axes @@ -850,6 +869,7 @@ def _drhs_run_tokamax_v2( num_actual_groups: int, rhs_dtype: jax.typing.DTypeLike, tiling: tuple, + use_gmm_v2_tiling_fn: bool, ) -> jnp.ndarray: """Executes Tokamax TGMM V2 backend for DRHS = LHS^T @ DRHS_dout.""" drhs_rhs = drhs_dout.qvalue if isinstance(drhs_dout, qpl.QArray) else drhs_dout @@ -859,7 +879,11 @@ def _drhs_run_tokamax_v2( if isinstance(drhs_dout, qpl.QArray): rhs_scale = _drhs_prepare_bwd_scale(drhs_dout) - custom_drhs_tiling = gmm_v2.TileSizes(tile_m=tiling[6], tile_k=tiling[7], tile_n=tiling[8]) + # Heuristic to calculate TGMM V2 tiling + drhs_tiling = tgmm_v2.calculate_tgmm_tiling + if not use_gmm_v2_tiling_fn: + # Custom tiling for DRHS + drhs_tiling = gmm_v2.TileSizes(tile_m=tiling[6], tile_k=tiling[7], tile_n=tiling[8]) return tgmm_v2.tgmm_v2( lhs=drhs_lhs, @@ -870,7 +894,7 @@ def _drhs_run_tokamax_v2( precision=jax.lax.Precision.DEFAULT, preferred_element_type=rhs_dtype, # pyrefly: ignore[bad-argument-type] group_offset=group_offset, - tile_info=custom_drhs_tiling, + tile_info=drhs_tiling, ) diff --git a/src/maxtext/layers/moe.py b/src/maxtext/layers/moe.py index da9e86e320..94b3d5adc7 100644 --- a/src/maxtext/layers/moe.py +++ b/src/maxtext/layers/moe.py @@ -1559,6 +1559,7 @@ def extract_vma(tensor): lhs_vma_axes=lhs_vma_axes, rhs_vma_axes=rhs_vma_axes, use_gmm_v2=self.config.use_gmm_v2, + use_gmm_v2_tiling_fn=self.config.use_gmm_v2_tiling_fn, partial_sum=partial_sum, interpret=megablox_interpret, ) diff --git a/tests/integration/tokamax_test.py b/tests/integration/tokamax_test.py index 784dc10607..1879538365 100644 --- a/tests/integration/tokamax_test.py +++ b/tests/integration/tokamax_test.py @@ -39,14 +39,15 @@ class Train(parameterized.TestCase): "use_gmm_v2": use_gmm_v2, "ici_expert_parallelism": ici_expert_parallelism, } - for base_name, quantization, use_gmm_v2, ici_expert_parallelism in [ - ("tokamax_v1_bf16", "", False, 1), - ("tokamax_v1_fp8", "fp8", False, 1), # not quantize gmm - ("tokamax_v1_fp8_full", "fp8_full", False, 1), # quantize gmm - ("tokamax_v2_bf16", "", True, 1), - ("tokamax_v2_fp8_full", "fp8_full", True, 1), - ("tokamax_v2_bf16", "", True, 2), - ("tokamax_v2_fp8_full", "fp8_full", True, 2), + for base_name, quantization, use_gmm_v2, use_gmm_v2_tiling_fn, ici_expert_parallelism in [ + ("tokamax_v1_bf16", "", False, False, 1), + ("tokamax_v1_fp8", "fp8", False, False, 1), # not quantize gmm + ("tokamax_v1_fp8_full", "fp8_full", False, False, 1), # quantize gmm + ("tokamax_v2_bf16", "", True, False, 1), + ("tokamax_v2_bf16_heuristic", "", True, True, 1), + ("tokamax_v2_fp8_full", "fp8_full", True, False, 1), + ("tokamax_v2_bf16", "", True, False, 2), + ("tokamax_v2_fp8_full", "fp8_full", True, False, 2), ] ) @pytest.mark.tpu_only @@ -54,6 +55,7 @@ def test_smoke_train( self, quantization: str, use_gmm_v2: bool, + use_gmm_v2_tiling_fn: bool, ici_expert_parallelism: int, ): """Smoke train with small config.""" @@ -84,6 +86,7 @@ def test_smoke_train( "megablox=False", "use_tokamax_gmm=True", f"use_gmm_v2={use_gmm_v2}", + f"use_gmm_v2_tiling_fn={use_gmm_v2_tiling_fn}", # tile sizes "wi_tile_fwd_batch_seq=128", "wi_tile_fwd_embed_dim=128",