From e9930f8240876d6ef96573cba3bd2a190be2ec42 Mon Sep 17 00:00:00 2001 From: Overwatch Agent Date: Fri, 7 Aug 2026 17:50:46 +0000 Subject: [PATCH 1/2] Fix DeepSeek4HyperHead compilation error --- src/maxtext/layers/mhc.py | 44 ++++++++++++++++++++++++++++-- src/maxtext/layers/nnx_decoders.py | 16 +++++++++-- 2 files changed, 55 insertions(+), 5 deletions(-) diff --git a/src/maxtext/layers/mhc.py b/src/maxtext/layers/mhc.py index 93b172dfa1..03e5f8aab0 100644 --- a/src/maxtext/layers/mhc.py +++ b/src/maxtext/layers/mhc.py @@ -24,8 +24,8 @@ from jax.sharding import Mesh from maxtext.common.common_types import Array, Config from maxtext.common.common_types import HyperConnectionType -from maxtext.layers.initializers import default_bias_init, default_scalar_init, nd_dense_init, variable_to_logically_partitioned -from maxtext.layers import nnx_wrappers +from maxtext.layers.initializers import default_bias_init, default_scalar_init, nd_dense_init +from maxtext.layers import linears from maxtext.layers.normalizations import RMSNorm @@ -313,4 +313,44 @@ def __call__( return res_out + post_out, metadata +class DeepSeek4HyperHead(nnx.Module): + """DeepSeek V4 Hyper Head.""" + def __init__( + self, + config: Config, + mesh: Mesh, + rngs: nnx.Rngs, + ): + self.config = config + self.mesh = mesh + self.rngs = rngs + self.k = config.mhc_expansion_rate + self.dim = config.emb_dim + self.dtype = config.dtype + self.weight_dtype = config.weight_dtype + + # tid2eid layers + self.tid2eid = nnx.List( + [ + linears.DenseGeneral( + in_features_shape=self.dim, + out_features_shape=self.dim, + dtype=self.dtype, + weight_dtype=self.weight_dtype, + rngs=self.rngs, + ) + for _ in range(config.first_num_hash_layers) + ] + ) + + def __call__(self, x: Array) -> Array: + # x shape: [batch, seq, expansion_rate, emb] + # Reduce expansion_rate dimension + x = jnp.sum(x, axis=2, dtype=x.dtype) + + # Apply tid2eid layers + for layer in self.tid2eid: + x = layer(x) + + return x diff --git a/src/maxtext/layers/nnx_decoders.py b/src/maxtext/layers/nnx_decoders.py index e6c0fc6f8a..895ea27c14 100644 --- a/src/maxtext/layers/nnx_decoders.py +++ b/src/maxtext/layers/nnx_decoders.py @@ -438,6 +438,13 @@ def __init__( self.is_gemma4 = self.config.decoder_block == DecoderBlockType.GEMMA4 self.is_gemma4_small = self.config.decoder_block == DecoderBlockType.GEMMA4_SMALL + if config.mhc_expansion_rate > 1 and config.decoder_block == DecoderBlockType.DEEPSEEK4: + self.hc_head = mhc.DeepSeek4HyperHead( + config=config, + mesh=self.mesh, + rngs=self.rngs, + ) + self._init_decoder_layers(decoder_block_classes, rngs, mesh) def _init_decoder_layers(self, decoder_block_classes, rngs, mesh): @@ -1967,13 +1974,16 @@ 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. # After the final transformer layer, `y` holds the raw, un-normalized hidden state. if getattr(cfg, "mhc_expansion_rate", 1) > 1: - # (batch, length, mhc_expansion_rate, emb_dim) --> (batch, length, emb_dim) - hidden_state = mhc_reduce(y) + if cfg.decoder_block == DecoderBlockType.DEEPSEEK4: + hidden_state = self.hc_head(y) + else: + # (batch, length, mhc_expansion_rate, emb_dim) --> (batch, length, emb_dim) + hidden_state = mhc_reduce(y) else: hidden_state = y + # When invoking from vLLM with RPA attention, logit computation is deferred to a later stage. if cfg.attention in ("vllm_rpa", "vllm_batched_rpa"): logits = None From f5a96305e5ea97a599ddb5d43a73a509a358b0a8 Mon Sep 17 00:00:00 2001 From: Overwatch Agent Date: Fri, 7 Aug 2026 17:53:02 +0000 Subject: [PATCH 2/2] Fix DeepSeek4HyperHead compilation error using nnx.Sequential --- src/maxtext/layers/mhc.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/src/maxtext/layers/mhc.py b/src/maxtext/layers/mhc.py index 03e5f8aab0..a1c1fbe912 100644 --- a/src/maxtext/layers/mhc.py +++ b/src/maxtext/layers/mhc.py @@ -331,8 +331,8 @@ def __init__( self.weight_dtype = config.weight_dtype # tid2eid layers - self.tid2eid = nnx.List( - [ + self.tid2eid = nnx.Sequential( + *[ linears.DenseGeneral( in_features_shape=self.dim, out_features_shape=self.dim, @@ -350,7 +350,6 @@ def __call__(self, x: Array) -> Array: x = jnp.sum(x, axis=2, dtype=x.dtype) # Apply tid2eid layers - for layer in self.tid2eid: - x = layer(x) + x = self.tid2eid(x) return x