diff --git a/src/maxtext/experimental/agent/ckpt_validation_pipeline/agent_sidecar/Dockerfile b/src/maxtext/experimental/agent/ckpt_validation_pipeline/agent_sidecar/Dockerfile index be88f5847b..a617319967 100644 --- a/src/maxtext/experimental/agent/ckpt_validation_pipeline/agent_sidecar/Dockerfile +++ b/src/maxtext/experimental/agent/ckpt_validation_pipeline/agent_sidecar/Dockerfile @@ -7,7 +7,9 @@ RUN apt-get update && apt-get install -y git curl ca-certificates gpg && \ curl -fsSL https://cli.github.com/packages/githubcli-archive-keyring.gpg | dd of=/usr/share/keyrings/githubcli-archive-keyring.gpg && \ chmod go+r /usr/share/keyrings/githubcli-archive-keyring.gpg && \ echo "deb [arch=$(dpkg --print-architecture) signed-by=/usr/share/keyrings/githubcli-archive-keyring.gpg] https://cli.github.com/packages stable main" | tee /etc/apt/sources.list.d/github-cli.list > /dev/null && \ - apt-get update && apt-get install -y gh && \ + echo "deb [signed-by=/usr/share/keyrings/cloud.google.gpg] http://packages.cloud.google.com/apt cloud-sdk main" | tee -a /etc/apt/sources.list.d/google-cloud-sdk.list && \ + curl https://packages.cloud.google.com/apt/doc/apt-key.gpg | gpg --dearmor -o /usr/share/keyrings/cloud.google.gpg && \ + apt-get update && apt-get install -y gh google-cloud-cli && \ rm -rf /var/lib/apt/lists/* # Copy only requirements first to leverage Docker cache for heavy installations diff --git a/src/maxtext/experimental/agent/ckpt_validation_pipeline/agent_sidecar/adk_agent.py b/src/maxtext/experimental/agent/ckpt_validation_pipeline/agent_sidecar/adk_agent.py index 9fea2794a6..a8c1f0701a 100644 --- a/src/maxtext/experimental/agent/ckpt_validation_pipeline/agent_sidecar/adk_agent.py +++ b/src/maxtext/experimental/agent/ckpt_validation_pipeline/agent_sidecar/adk_agent.py @@ -11,7 +11,7 @@ logger = logging.getLogger(__name__) -def _send_message_with_retry(chat, prompt, max_retries=3, sleep_seconds=30): +def _send_message_with_retry(chat, prompt, max_retries=5, sleep_seconds=60): """Sends a message to Gemini with retry and a 30-second sleep on 429 rate-limit/quota errors.""" for attempt in range(1, max_retries + 1): try: diff --git a/src/maxtext/layers/mhc.py b/src/maxtext/layers/mhc.py index a1c1fbe912..362bf2c9be 100644 --- a/src/maxtext/layers/mhc.py +++ b/src/maxtext/layers/mhc.py @@ -25,7 +25,6 @@ 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 -from maxtext.layers import linears from maxtext.layers.normalizations import RMSNorm @@ -314,7 +313,7 @@ def __call__( class DeepSeek4HyperHead(nnx.Module): - """DeepSeek V4 Hyper Head.""" + """Implements DeepSeek4 HyperHead.""" def __init__( self, @@ -323,33 +322,44 @@ def __init__( rngs: nnx.Rngs, ): self.config = config - self.mesh = mesh + self.hc_mult = config.mhc_expansion_rate 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.Sequential( - *[ - 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) - ] + self.mesh = mesh + self.dtype = self.config.dtype + self.weight_dtype = self.config.weight_dtype + self.eps = 1e-6 + + self.input_norm = RMSNorm( + num_features=self.hc_mult * config.emb_dim, + dtype=self.dtype, + weight_dtype=self.weight_dtype, + kernel_axes=("norm",), + epsilon=config.normalization_layer_epsilon, + rngs=self.rngs, + ) + + self.hc_fn = nnx.Param( + default_scalar_init(self.rngs.params(), (self.hc_mult, self.hc_mult * config.emb_dim), self.weight_dtype), + out_sharding=(None, None), + ) + self.hc_base = nnx.Param( + default_scalar_init(self.rngs.params(), (self.hc_mult,), self.weight_dtype), + out_sharding=(None,), + ) + self.hc_scale = nnx.Param( + default_scalar_init(self.rngs.params(), (1,), self.weight_dtype), + out_sharding=(None,), ) 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) + b, s, k, d = x.shape + flat = jnp.reshape(x, (b, s, k * d)) + flat = self.input_norm(flat) - # Apply tid2eid layers - x = self.tid2eid(x) + mixes = jnp.einsum("bsm,nm->bsn", flat, jnp.asarray(self.hc_fn[...], self.dtype)) + pre = ( + jax.nn.sigmoid(mixes * jnp.asarray(self.hc_scale[...], self.dtype) + jnp.asarray(self.hc_base[...], self.dtype)) + + self.eps + ) - return x + return jnp.sum(x * jnp.expand_dims(pre, axis=3), axis=2)