Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 16 additions & 5 deletions src/maxtext/layers/moe.py
Original file line number Diff line number Diff line change
Expand Up @@ -3363,12 +3363,23 @@ def fused_moe_matmul(
self.config.decoder_block not in (ctypes.DecoderBlockType.LLAMA4, ctypes.DecoderBlockType.GEMMA4)
)

output_2d = fused_moe_func(
# The fused kernel quantizes for itself: expert weights go in pre-quantized with
# per-block scales (when the qwix rule covers the grouped matmul) and the kernel
# quantizes the activations in-kernel, accumulating in f32. The call runs outside
# qwix's interception: qwix only guards pallas_call, while the kernel is launched
# through pl.kernel, so under a QtProvider its tiled matmuls would otherwise be
# fake-quantized into fp8 x fp8 with a bf16 accumulator, which Mosaic rejects.
rule = quantizations.get_fused_moe_rule()
quantized_w1, w1_scale = quantizations.quantize_weight_for_fused_moe(fused_kernel, rule)
quantized_w2, w2_scale = quantizations.quantize_weight_for_fused_moe(wo_kernel, rule)
fused_moe = quantizations.without_qwix_interception(fused_moe_func)

output_2d = fused_moe(
hidden_states=hidden_states,
w1=fused_kernel,
w2=wo_kernel,
w1_scale=None,
w2_scale=None,
w1=quantized_w1,
w2=quantized_w2,
w1_scale=w1_scale,
w2_scale=w2_scale,
w1_bias=None,
w2_bias=None,
gating_output=gating_output,
Expand Down
70 changes: 70 additions & 0 deletions src/maxtext/layers/quantizations.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
from qwix._src.core import numerics
from qwix._src.core import dot_general_qt
from qwix._src.core import sparsity
from qwix._src import interception as qwix_interception

import jax
import jax.numpy as jnp
Expand Down Expand Up @@ -66,6 +67,7 @@ def _safe_find_param(x, ptq_array_type=None):
from maxtext.configs.types import TeCommGemmOverlapPolicy
from maxtext.common.common_types import DType, Config
from maxtext.inference.kvcache import KVQuant
from maxtext.utils import max_logging

# Params used to define mixed precision quantization configs
DEFAULT = "__default__" # default config
Expand Down Expand Up @@ -970,6 +972,74 @@ def maybe_quantize_model(model, config):
return model


# Quantized weight dtypes that tpu-inference's fused MoE kernel (gmm_v2) takes with
# per-block scales and dequantizes in-kernel. Any other qtype keeps the expert weights
# unquantized on the fused path.
FUSED_MOE_KERNEL_WEIGHT_QTYPES = (jnp.dtype(jnp.float8_e4m3fn), jnp.dtype(jnp.int8))


def get_fused_moe_rule() -> qwix.QuantizationRule | None:
"""Returns the qwix rule that governs the fused MoE grouped matmul, if any.

The fused kernel is the same grouped matmul that MaxText's own megablox / ragged_dot
paths implement, so it takes its rule from the "gmm" op exactly like they do. Rules
that only list "dot_general" (the plain fp8 / int8 recipes) leave the experts
unquantized on every MoE path, including this one.
"""
return qpl.get_current_rule("gmm")


def quantize_weight_for_fused_moe(
kernel: jax.Array, rule: qwix.QuantizationRule | None
) -> tuple[jax.Array, jax.Array | None]:
"""Quantizes an [E, K, N] expert weight for tpu-inference's fused MoE kernel.

The kernel contracts over K, dequantizes after the matmul with scales laid out as
[E, num_blocks, 1, N] (blocks along K), and quantizes the activations itself with a
32-bit accumulator. Scale granularity follows the qwix rule: channelwise over (E, N),
plus blockwise along K when the rule sets tile_size.

Returns the (possibly unchanged) weight and its scale, or None when the rule does not
ask for a weight dtype the kernel supports.
"""

weight_qtype = getattr(rule, "weight_qtype", None) if rule is not None else None
if weight_qtype is None:
return kernel, None
qtype = jnp.dtype(weight_qtype)
if qtype not in FUSED_MOE_KERNEL_WEIGHT_QTYPES:
max_logging.log(f"fused MoE kernel does not take {qtype} weights; keeping expert weights in {kernel.dtype}.")
return kernel, None
if kernel.ndim != 3:
raise ValueError(f"fused MoE expert weights must be [E, K, N], got {kernel.shape}")

tile_size = getattr(rule, "tile_size", None)
tiled_axes = {1: tile_size} if tile_size else {}
cal_method = getattr(rule, "weight_calibration_method", None)
quant_kwargs = {"calibration_method": cal_method} if cal_method is not None else {}

quantized = qpl.quantize(
kernel,
qtype,
channelwise_axes=(0, 2),
tiled_axes=tiled_axes,
scale_dtype=jnp.float32,
**quant_kwargs,
)
scale = jnp.expand_dims(quantized.scale, axis=2)
return quantized.qvalue, scale


def without_qwix_interception(fn: Callable) -> Callable:
"""Returns `fn` wrapped so that qwix's op interception is off while it runs.

For ops that do their own quantization (tpu-inference's fused MoE), the qwix rule is
applied once up front to their inputs and the op itself is opaque to qwix: neither
the routing math around the kernel nor anything traced inside it gets rewritten.
"""
return qwix_interception.disable_interceptions(fn)


def _cast_reduced_from(arr, reduced_arr):
aval = jax.typeof(reduced_arr)
# In shard map
Expand Down
264 changes: 264 additions & 0 deletions tests/unit/fused_moe_qwix_test.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,264 @@
# 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
#
# https://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.

# pylint: disable=missing-class-docstring,missing-function-docstring,unbalanced-tuple-unpacking

"""Tests for the qwix boundary around tpu-inference's fused MoE kernel."""

import sys
import types
import unittest
from unittest import mock

from flax import nnx
import jax
import jax.numpy as jnp
import numpy as np
import qwix

from maxtext.common import common_types as ctypes
from maxtext.layers import moe
from maxtext.layers import quantizations


def _assert_error_within(actual, ref, bound):
"""Every |actual - ref| stays under `bound`, which broadcasts against the arrays."""
err = np.abs(np.asarray(actual - ref))
np.testing.assert_array_less(err, np.broadcast_to(np.asarray(bound) + 1e-6, err.shape))


def _fp8_rule(**kwargs):
return qwix.QtRule(
module_path=".*",
weight_qtype=jnp.float8_e4m3fn,
act_qtype=jnp.float8_e4m3fn,
op_names=("dot_general", "gmm", "ragged_dot"),
**kwargs,
)


class QuantizeWeightForFusedMoeTest(unittest.TestCase):
"""quantize_weight_for_fused_moe produces what the fused kernel takes."""

def setUp(self):
super().setUp()
self.num_experts, self.in_dim, self.out_dim = 4, 256, 128
self.kernel = jax.random.normal(jax.random.key(0), (self.num_experts, self.in_dim, self.out_dim), jnp.bfloat16)

def test_no_rule_keeps_weight(self):
weight, scale = quantizations.quantize_weight_for_fused_moe(self.kernel, None)
self.assertIs(weight, self.kernel)
self.assertIsNone(scale)

def test_rule_without_weight_qtype_keeps_weight(self):
rule = qwix.QtRule(module_path=".*", act_qtype=jnp.float8_e4m3fn)
weight, scale = quantizations.quantize_weight_for_fused_moe(self.kernel, rule)
self.assertIs(weight, self.kernel)
self.assertIsNone(scale)

def test_unsupported_qtype_keeps_weight(self):
rule = qwix.QtRule(module_path=".*", weight_qtype=jnp.float8_e5m2)
weight, scale = quantizations.quantize_weight_for_fused_moe(self.kernel, rule)
self.assertIs(weight, self.kernel)
self.assertIsNone(scale)

def test_fp8_channelwise(self):
weight, scale = quantizations.quantize_weight_for_fused_moe(self.kernel, _fp8_rule())
self.assertEqual(weight.dtype, jnp.float8_e4m3fn)
self.assertEqual(weight.shape, self.kernel.shape)
# one scale per (expert, output column), laid out as the kernel expects
self.assertEqual(scale.shape, (self.num_experts, 1, 1, self.out_dim))
self.assertEqual(scale.dtype, jnp.float32)
dequantized = weight.astype(jnp.float32) * scale[:, 0]
ref = self.kernel.astype(jnp.float32)
# fp8 e4m3 keeps 3 mantissa bits: the error is bounded by half an ulp of the column max
col_max = jnp.max(jnp.abs(ref), axis=1, keepdims=True)
_assert_error_within(dequantized, ref, col_max / 16)

def test_fp8_blockwise_from_tile_size(self):
weight, scale = quantizations.quantize_weight_for_fused_moe(self.kernel, _fp8_rule(tile_size=64))
self.assertEqual(weight.dtype, jnp.float8_e4m3fn)
self.assertEqual(scale.shape, (self.num_experts, self.in_dim // 64, 1, self.out_dim))
blocks = weight.astype(jnp.float32).reshape(self.num_experts, self.in_dim // 64, 64, self.out_dim)
dequantized = (blocks * scale).reshape(self.kernel.shape)
ref = self.kernel.astype(jnp.float32)
block_max = jnp.max(jnp.abs(ref.reshape(blocks.shape)), axis=2, keepdims=True)
_assert_error_within(dequantized.reshape(blocks.shape), ref.reshape(blocks.shape), block_max / 16)

def test_int8_channelwise(self):
rule = qwix.QtRule(module_path=".*", weight_qtype=jnp.int8, op_names=("gmm",))
weight, scale = quantizations.quantize_weight_for_fused_moe(self.kernel, rule)
self.assertEqual(weight.dtype, jnp.int8)
self.assertEqual(scale.shape, (self.num_experts, 1, 1, self.out_dim))
dequantized = weight.astype(jnp.float32) * scale[:, 0]
ref = self.kernel.astype(jnp.float32)
col_max = jnp.max(jnp.abs(ref), axis=1, keepdims=True)
_assert_error_within(dequantized, ref, col_max / 127)

def test_rejects_non_3d_weight(self):
with self.assertRaises(ValueError):
quantizations.quantize_weight_for_fused_moe(self.kernel[0], _fp8_rule())

def test_rule_with_missing_attributes(self):
"""getattr fallback handles rules missing tile_size or calibration_method."""

class MinimalRule:
weight_qtype = jnp.float8_e4m3fn

weight, scale = quantizations.quantize_weight_for_fused_moe(self.kernel, MinimalRule())
self.assertEqual(weight.dtype, jnp.float8_e4m3fn)
self.assertEqual(scale.shape, (self.num_experts, 1, 1, self.out_dim))


class WithoutQwixInterceptionTest(unittest.TestCase):
"""The boundary hides the qwix rule (and its op rewriting) from whatever runs inside."""

def test_rule_visible_outside_and_hidden_inside(self):
seen = {}

class Layer(nnx.Module):

def __init__(self):
self.w = nnx.Param(jnp.ones((32, 64), jnp.bfloat16))

def __call__(self, x):
seen["outside"] = quantizations.get_fused_moe_rule()

def opaque(y):
seen["inside"] = quantizations.get_fused_moe_rule()
return jnp.dot(y, self.w[...], preferred_element_type=jnp.float32)

out = quantizations.without_qwix_interception(opaque)(x)
seen["after"] = quantizations.get_fused_moe_rule()
return out

x = jax.random.normal(jax.random.key(0), (8, 32), jnp.bfloat16)
layer = qwix.quantize_model(Layer(), qwix.QtProvider([_fp8_rule()]), x)
jaxpr = str(jax.make_jaxpr(layer)(x))

self.assertIsNotNone(seen["outside"])
self.assertIsNone(seen["inside"])
self.assertIsNotNone(seen["after"])
# the dot inside the boundary is left alone: no fp8 casts anywhere in the trace
self.assertNotIn("float8", jaxpr)


class FusedMoeMatmulTest(unittest.TestCase):
"""RoutedMoE.fused_moe_matmul hands the kernel pre-quantized weights, outside qwix."""

num_experts, top_k, emb_dim, mlp_dim, tokens = 4, 2, 64, 128, 16

def _fake_tpu_inference(self, calls):
"""Stand-in for the tpu-inference modules fused_moe_matmul imports."""

def fused_moe_func(**kwargs):
calls.append(dict(kwargs, rule_inside=quantizations.get_fused_moe_rule()))
return jnp.zeros((kwargs["hidden_states"].shape[0], self.emb_dim), jnp.bfloat16)

envs = types.SimpleNamespace(
ENABLE_RS_KERNEL=False, USE_GMM_FUSED_RS_KERNEL=False, ONEHOT_MOE_PERMUTE_THRESHOLD=0, VLLM_MOE_CHUNK_SIZE=0
)
pkg = types.ModuleType("tpu_inference")
pkg.envs = envs
layers = types.ModuleType("tpu_inference.layers")
common = types.ModuleType("tpu_inference.layers.common")
gmm = types.ModuleType("tpu_inference.layers.common.fused_moe_gmm")
gmm.fused_moe_func = fused_moe_func
return {
"tpu_inference": pkg,
"tpu_inference.envs": envs,
"tpu_inference.layers": layers,
"tpu_inference.layers.common": common,
"tpu_inference.layers.common.fused_moe_gmm": gmm,
}

def _make_layer(self, calls):
test = self
config = types.SimpleNamespace(
mlp_activations=("silu",),
routed_score_func="softmax",
norm_topk_prob=True,
decoder_block=ctypes.DecoderBlockType.MIXTRAL,
)

class Layer(nnx.Module):
"""Drives RoutedMoE.fused_moe_matmul with just the attributes it reads."""

def __init__(self):
keys = jax.random.split(jax.random.key(0), 3)
self.w0 = nnx.Param(jax.random.normal(keys[0], (test.num_experts, test.emb_dim, test.mlp_dim), jnp.bfloat16))
self.w1 = nnx.Param(jax.random.normal(keys[1], (test.num_experts, test.emb_dim, test.mlp_dim), jnp.bfloat16))
self.wo = nnx.Param(jax.random.normal(keys[2], (test.num_experts, test.mlp_dim, test.emb_dim), jnp.bfloat16))
self.config = config
self.num_experts = test.num_experts
self.num_experts_per_tok = test.top_k
self.mesh = None

def get_expert_parallelism_size(self):
return 1

def __call__(self, inputs, gate_logits):
with mock.patch.dict(sys.modules, test._fake_tpu_inference(calls)):
out, _, _ = moe.RoutedMoE.fused_moe_matmul(self, inputs, gate_logits, self.wo[...], self.w0[...], self.w1[...])
return out

return Layer()

def _inputs(self):
inputs = jax.random.normal(jax.random.key(1), (1, self.tokens, self.emb_dim), jnp.bfloat16)
gate_logits = jax.random.normal(jax.random.key(2), (1, self.tokens, self.num_experts), jnp.float32)
return inputs, gate_logits

def test_unquantized_without_qwix(self):
calls = []
inputs, gate_logits = self._inputs()
self._make_layer(calls)(inputs, gate_logits)
(call,) = calls
self.assertEqual(call["w1"].dtype, jnp.bfloat16)
self.assertEqual(call["w1"].shape, (self.num_experts, self.emb_dim, 2 * self.mlp_dim))
self.assertIsNone(call["w1_scale"])
self.assertIsNone(call["w2_scale"])

def test_fp8_rule_prequantizes_weights_outside_qwix(self):
calls = []
inputs, gate_logits = self._inputs()
layer = qwix.quantize_model(self._make_layer(calls), qwix.QtProvider([_fp8_rule()]), inputs, gate_logits)
calls.clear()
layer(inputs, gate_logits)
(call,) = calls
# weights arrive quantized, with scales in the kernel's [E, blocks, 1, N] layout
self.assertEqual(call["w1"].dtype, jnp.float8_e4m3fn)
self.assertEqual(call["w2"].dtype, jnp.float8_e4m3fn)
self.assertEqual(call["w1_scale"].shape, (self.num_experts, 1, 1, 2 * self.mlp_dim))
self.assertEqual(call["w2_scale"].shape, (self.num_experts, 1, 1, self.emb_dim))
# and the kernel itself runs outside qwix's interception
self.assertIsNone(call["rule_inside"])

def test_dot_general_only_rule_keeps_experts_unquantized(self):
calls = []
inputs, gate_logits = self._inputs()
rule = qwix.QtRule(
module_path=".*", weight_qtype=jnp.float8_e4m3fn, act_qtype=jnp.float8_e4m3fn, op_names=("dot_general",)
)
layer = qwix.quantize_model(self._make_layer(calls), qwix.QtProvider([rule]), inputs, gate_logits)
calls.clear()
layer(inputs, gate_logits)
(call,) = calls
self.assertEqual(call["w1"].dtype, jnp.bfloat16)
self.assertIsNone(call["w1_scale"])
self.assertIsNone(call["rule_inside"])


if __name__ == "__main__":
unittest.main()
Loading