From 14a680e4f28f0ac12ced049ad3e08cd518bb8bda Mon Sep 17 00:00:00 2001 From: Denis Milovanov Date: Wed, 9 Sep 2026 09:47:04 -0700 Subject: [PATCH] Fix: Numeric divergence between the mHC layer and its Pallas kernel --- src/maxtext/kernels/mhc/__init__.py | 2 + src/maxtext/kernels/mhc/api.py | 1 + src/maxtext/layers/mhc.py | 26 ++-- tests/unit/mhc_test.py | 233 ++++++++++++++++------------ 4 files changed, 150 insertions(+), 112 deletions(-) diff --git a/src/maxtext/kernels/mhc/__init__.py b/src/maxtext/kernels/mhc/__init__.py index 3a36d44e16..94af289d6f 100644 --- a/src/maxtext/kernels/mhc/__init__.py +++ b/src/maxtext/kernels/mhc/__init__.py @@ -13,6 +13,7 @@ # limitations under the License. """MaxText mHC-lite Pallas kernel package.""" +from maxtext.kernels.mhc.api import compute_sigmoid_gate from maxtext.kernels.mhc.api import hbm_specs from maxtext.kernels.mhc.api import MhcCoeffGradients from maxtext.kernels.mhc.api import MhcCoeffOutputs @@ -37,4 +38,5 @@ "MhcCoeffGradients", "UnsupportedInputError", "hbm_specs", + "compute_sigmoid_gate", ] diff --git a/src/maxtext/kernels/mhc/api.py b/src/maxtext/kernels/mhc/api.py index f1f0b77c4c..a821c52d31 100644 --- a/src/maxtext/kernels/mhc/api.py +++ b/src/maxtext/kernels/mhc/api.py @@ -27,6 +27,7 @@ MhcCoeffOutputs = common.MhcCoeffOutputs MhcCoeffGradients = common.MhcCoeffGradients hbm_specs = common.hbm_specs +compute_sigmoid_gate = common.compute_sigmoid_gate def _validate_implementation( diff --git a/src/maxtext/layers/mhc.py b/src/maxtext/layers/mhc.py index 81f93b58fa..ae3d07eafd 100644 --- a/src/maxtext/layers/mhc.py +++ b/src/maxtext/layers/mhc.py @@ -234,15 +234,6 @@ def res_mapping(self, h_res: Array): output = sinkhorn(intermediate, self.sinkhorn_iterations) return output - def mapping(self, h: Array, alpha_scale: Array, beta: Array, scale: float, eps: float = 0.0): - """Helper function for both pre and post mappings after matmul.""" - # In MaxText, we match weight precision to activations before Matmul - beta = jnp.asarray(beta, self.dtype) - alpha_scale = jnp.asarray(alpha_scale, self.dtype) - intermediate = alpha_scale * h + beta[None, None, :] - output = scale * jax.nn.sigmoid(intermediate) + eps - return output - def __call__( self, norm_fn: Callable, @@ -306,13 +297,15 @@ def __call__( h_res = h_concat[..., 2 * self.k :] # 2. Pre mapping - pre_mapping = self.mapping( + # Shared with the Pallas kernel so both paths gate identically. The + # helper computes in float32; cast back to keep the GEMM in self.dtype. + pre_mapping = mhc_kernel.compute_sigmoid_gate( h_pre, self.pre_alpha_scale[...], self.pre_beta[...], - 1.0, - eps=1e-6, - ) + multiplier=1.0, + epsilon=1e-6, + ).astype(self.dtype) # bskd, bsk -> bsd (fused contracted GEMM) layer_input = jnp.einsum( "bsk,bskd->bsd", @@ -354,12 +347,13 @@ def __call__( return output, metadata # 5. Post mapping - post_mapping = self.mapping( + post_mapping = mhc_kernel.compute_sigmoid_gate( h_post, self.post_alpha_scale[...], self.post_beta[...], - 2.0, - ) + multiplier=2.0, + epsilon=0.0, + ).astype(self.dtype) # Moving away from einsum seems to allow XLA to perform better fusions # bsd,bsk -> bskd post_out = jnp.expand_dims(layer_out, axis=2) * jnp.expand_dims(post_mapping, axis=3) diff --git a/tests/unit/mhc_test.py b/tests/unit/mhc_test.py index db0cd49c85..85be85517d 100644 --- a/tests/unit/mhc_test.py +++ b/tests/unit/mhc_test.py @@ -14,6 +14,7 @@ """Test for DeepSeek Manifold-Constrained Hyper Connections (mHC).""" +import contextlib import dataclasses import itertools import math @@ -93,6 +94,40 @@ def test_doubly_stochastic_property(self): np.testing.assert_allclose(col_sums, jnp.ones_like(col_sums), atol=1e-3) +@contextlib.contextmanager +def _interpreted_mhc_kernel(interpret: bool = True): + """Forces the mHC Pallas kernel into interpret mode so it executes off-TPU. + + Interpret mode traces the kernel body as plain JAX and emulates the grid in + Python, so it validates the algorithm but never exercises Mosaic lowering, + the real VMEM budget, or MXU accumulation. Pass ``interpret=False`` to leave + the kernel entry points untouched, that path needs TPU hardware. + + Yields the patched ``(pre, post)`` when interpret=True + or ``(None, None)`` when interpret=False. + """ + if not interpret: + yield None, None + return + + real_pre = mhc.mhc_kernel.pre + real_post = mhc.mhc_kernel.post + + def force_interpret(fn): + def wrapped(*args, **kwargs): + config = kwargs.get("config", mhc.mhc_kernel.MhcKernelConfig()) + kwargs["config"] = dataclasses.replace(config, interpret=True) + return fn(*args, **kwargs) + + return wrapped + + with ( + mock.patch.object(mhc.mhc_kernel, "pre", side_effect=force_interpret(real_pre)) as mock_pre, + mock.patch.object(mhc.mhc_kernel, "post", side_effect=force_interpret(real_post)) as mock_post, + ): + yield mock_pre, mock_post + + class TestMHC(parameterized.TestCase): """Test for MHC module""" @@ -175,6 +210,23 @@ def _setup_mhc( rngs=self.rngs, ) + def _build_mhc_and_mlp(self): + """Builds the mHC module and the MLP branch it wraps, for the active config.""" + module = mhc.ManifoldConstrainedHyperConnections(self.config, self.dim, self.mesh, self.rngs) + layer = linears.MlpBlock( + config=self.config, + mesh=self.mesh, + in_features=self.config.emb_dim, + intermediate_dim=self.config.moe_mlp_dim, + activations=self.config.mlp_activations, + intermediate_dropout_rate=self.config.dropout_rate, + dtype=self.config.dtype, + weight_dtype=self.config.weight_dtype, + model_mode=self.config.model_call_mode, + rngs=self.rngs, + ) + return module, layer + # Skip GPU due to NotImplementedError: dynamic grid bounds not supported in the Triton backend @pytest.mark.tpu_only @parameterized.named_parameters(("Rate3", 3), ("Rate4", 4)) @@ -208,19 +260,7 @@ def test_moe_layer_output_shape(self, rate): def test_dense_layer_output_shape(self, rate): self._setup_mhc(rate) with nn_partitioning.axis_rules(self.config.logical_axis_rules): - module = mhc.ManifoldConstrainedHyperConnections(self.config, self.dim, self.mesh, self.rngs) - layer = linears.MlpBlock( - config=self.config, - mesh=self.mesh, - in_features=self.config.emb_dim, - intermediate_dim=self.config.moe_mlp_dim, - activations=self.config.mlp_activations, - intermediate_dropout_rate=self.config.dropout_rate, - dtype=self.config.dtype, - weight_dtype=self.config.weight_dtype, - model_mode=self.config.model_call_mode, - rngs=self.rngs, - ) + module, layer = self._build_mhc_and_mlp() b, s, k, d = self.x.shape output, metadata = module(self.pre_norm, layer, x=self.x, mhc_type=HyperConnectionType.MLP_DENSE) @@ -408,39 +448,9 @@ def test_use_mhc_pallas_kernel_dispatch(self, use_mhc_pallas_kernel): dtype="bfloat16", ) with nn_partitioning.axis_rules(self.config.logical_axis_rules): - module = mhc.ManifoldConstrainedHyperConnections(self.config, self.dim, self.mesh, self.rngs) - layer = linears.MlpBlock( - config=self.config, - mesh=self.mesh, - in_features=self.config.emb_dim, - intermediate_dim=self.config.moe_mlp_dim, - activations=self.config.mlp_activations, - intermediate_dropout_rate=self.config.dropout_rate, - dtype=self.config.dtype, - weight_dtype=self.config.weight_dtype, - model_mode=self.config.model_call_mode, - rngs=self.rngs, - ) + module, layer = self._build_mhc_and_mlp() - real_pre = mhc.mhc_kernel.pre - real_post = mhc.mhc_kernel.post - - def fake_pre(*args, **kwargs): - config = kwargs.get("config", mhc.mhc_kernel.MhcKernelConfig()) - config = dataclasses.replace(config, interpret=True) - kwargs["config"] = config - return real_pre(*args, **kwargs) - - def fake_post(*args, **kwargs): - config = kwargs.get("config", mhc.mhc_kernel.MhcKernelConfig()) - config = dataclasses.replace(config, interpret=True) - kwargs["config"] = config - return real_post(*args, **kwargs) - - with ( - mock.patch.object(mhc.mhc_kernel, "pre", side_effect=fake_pre) as mock_pre, - mock.patch.object(mhc.mhc_kernel, "post", side_effect=fake_post) as mock_post, - ): + with _interpreted_mhc_kernel() as (mock_pre, mock_post): output, _ = module( self.pre_norm, layer, @@ -470,6 +480,85 @@ def fake_post(*args, **kwargs): self.assertEqual(output.shape, self.x.shape) + def _assert_pallas_matches_reference(self, *, interpret): + """Asserts the kernel path matches the reference path, in interpret mode or on hardware.""" + setup_kwargs = { + "enable_mhc_lite": True, + "dim": 128, + "sequence_length": 256, + "per_device_batch_size": 1, + "dtype": "bfloat16", + } + + self._setup_mhc(4, use_mhc_pallas_kernel=False, **setup_kwargs) + with nn_partitioning.axis_rules(self.config.logical_axis_rules): + module, layer = self._build_mhc_and_mlp() + shared_state = jax.tree.map( + jnp.copy, + (nnx.state(module), nnx.state(layer), nnx.state(self.pre_norm)), + ) + reference_output, _ = module( + self.pre_norm, + layer, + x=self.x, + mhc_type=HyperConnectionType.MLP_DENSE, + ) + + self._setup_mhc(4, use_mhc_pallas_kernel=True, **setup_kwargs) + with nn_partitioning.axis_rules(self.config.logical_axis_rules): + module, layer = self._build_mhc_and_mlp() + module_state, layer_state, norm_state = shared_state + nnx.update(module, module_state) + nnx.update(layer, layer_state) + nnx.update(self.pre_norm, norm_state) + with _interpreted_mhc_kernel(interpret): + kernel_output, _ = module( + self.pre_norm, + layer, + x=self.x, + mhc_type=HyperConnectionType.MLP_DENSE, + ) + + self.assertEqual(kernel_output.dtype, reference_output.dtype) + # Both branches emit bfloat16, whose spacing at this output magnitude is + # ~0.03. The tolerance is that noise floor: this test guards against + # structural divergence (wrong constant, epsilon, or slice) between the two + # branches, not against sub-ULP differences in accumulation order. + np.testing.assert_allclose( + kernel_output.astype(jnp.float32), + reference_output.astype(jnp.float32), + rtol=5e-2, + atol=5e-2, + ) + + def test_pallas_kernel_matches_reference_path(self): + """Toggling use_mhc_pallas_kernel must not change the layer's output values.""" + self._assert_pallas_matches_reference(interpret=True) + + @pytest.mark.tpu_only + def test_pallas_kernel_matches_reference_path_tpu(self): + """As above, but against the real Mosaic-compiled kernel instead of interpret mode.""" + self._assert_pallas_matches_reference(interpret=False) + + def test_sigmoid_gate_computes_in_float32(self): + """The shared gate must not round its scale/bias down to the activation dtype. + + Both the layer and the kernel call `compute_sigmoid_gate`, so a bfloat16 + regression here would silently re-introduce a numerical difference between + the two branches that the end-to-end comparison above cannot resolve. + """ + key = jax.random.PRNGKey(0) + logits = jax.random.normal(key, (8, 4), dtype=jnp.bfloat16) + # Values that are not representable in bfloat16 so rounding is observable. + scale = jnp.asarray([1.0001, 0.9999, 1.0002, 0.9998], dtype=jnp.float32) + bias = jnp.asarray([0.0001, -0.0001, 0.0002, -0.0002], dtype=jnp.float32) + + gate = mhc_kernel.compute_sigmoid_gate(logits, scale, bias, multiplier=2.0, epsilon=1e-6) + + expected = 2.0 * jax.nn.sigmoid(scale * logits.astype(jnp.float32) + bias) + 1e-6 + self.assertEqual(gate.dtype, jnp.float32) + np.testing.assert_allclose(gate, expected, rtol=1e-6, atol=1e-6) + def test_use_mhc_pallas_kernel_custom_block_size(self): """Verify that custom block sizes are passed to the kernel.""" self._setup_mhc( @@ -485,39 +574,9 @@ def test_use_mhc_pallas_kernel_custom_block_size(self): dtype="bfloat16", ) with nn_partitioning.axis_rules(self.config.logical_axis_rules): - module = mhc.ManifoldConstrainedHyperConnections(self.config, self.dim, self.mesh, self.rngs) - layer = linears.MlpBlock( - config=self.config, - mesh=self.mesh, - in_features=self.config.emb_dim, - intermediate_dim=self.config.moe_mlp_dim, - activations=self.config.mlp_activations, - intermediate_dropout_rate=self.config.dropout_rate, - dtype=self.config.dtype, - weight_dtype=self.config.weight_dtype, - model_mode=self.config.model_call_mode, - rngs=self.rngs, - ) + module, layer = self._build_mhc_and_mlp() - real_pre = mhc.mhc_kernel.pre - real_post = mhc.mhc_kernel.post - - def fake_pre(*args, **kwargs): - config = kwargs.get("config", mhc.mhc_kernel.MhcKernelConfig()) - config = dataclasses.replace(config, interpret=True) - kwargs["config"] = config - return real_pre(*args, **kwargs) - - def fake_post(*args, **kwargs): - config = kwargs.get("config", mhc.mhc_kernel.MhcKernelConfig()) - config = dataclasses.replace(config, interpret=True) - kwargs["config"] = config - return real_post(*args, **kwargs) - - with ( - mock.patch.object(mhc.mhc_kernel, "pre", side_effect=fake_pre) as mock_pre, - mock.patch.object(mhc.mhc_kernel, "post", side_effect=fake_post) as mock_post, - ): + with _interpreted_mhc_kernel() as (mock_pre, mock_post): output, _ = module( self.pre_norm, layer, @@ -604,25 +663,7 @@ def forward_baseline(x): ) module_kernel = mhc.ManifoldConstrainedHyperConnections(self.config, self.dim, self.mesh, self.rngs) - real_pre = mhc.mhc_kernel.pre - real_post = mhc.mhc_kernel.post - - def fake_pre(*args, **kwargs): - cfg = kwargs.get("config", mhc.mhc_kernel.MhcKernelConfig()) - cfg = dataclasses.replace(cfg, interpret=True) - kwargs["config"] = cfg - return real_pre(*args, **kwargs) - - def fake_post(*args, **kwargs): - cfg = kwargs.get("config", mhc.mhc_kernel.MhcKernelConfig()) - cfg = dataclasses.replace(cfg, interpret=True) - kwargs["config"] = cfg - return real_post(*args, **kwargs) - - with ( - mock.patch.object(mhc.mhc_kernel, "pre", side_effect=fake_pre), - mock.patch.object(mhc.mhc_kernel, "post", side_effect=fake_post), - ): + with _interpreted_mhc_kernel(): def forward_kernel(x): out, _ = module_kernel(self.pre_norm, layer_fn, x=x, mhc_type=HyperConnectionType.MLP_DENSE)