diff --git a/src/dependencies/requirements/requirements.txt b/src/dependencies/requirements/requirements.txt index 05c2be074b..95a38ecbd4 100644 --- a/src/dependencies/requirements/requirements.txt +++ b/src/dependencies/requirements/requirements.txt @@ -33,13 +33,13 @@ pylint pytest pytype qwix>=0.1.6 -sentencepiece +sentencepiece>=0.2.0 tensorboard-plugin-profile tensorboardx tensorflow-datasets tensorflow-text tensorflow -tiktoken +tiktoken>=0.5.0 tokamax>=0.0.4 transformers google-jetstream @ https://github.com/AI-Hypercomputer/JetStream/archive/29329e8e73820993f77cfc8efe34eb2a73f5de98.zip diff --git a/src/maxtext/checkpoint_conversion/to_maxtext.py b/src/maxtext/checkpoint_conversion/to_maxtext.py index ce2eecb678..4ba4cf9eac 100644 --- a/src/maxtext/checkpoint_conversion/to_maxtext.py +++ b/src/maxtext/checkpoint_conversion/to_maxtext.py @@ -53,6 +53,7 @@ from functools import partial import json import os +import re import sys import threading import time @@ -119,20 +120,25 @@ def __init__(self, model_id, token, revision=None): self.current_shard_content = {} # Cache for resolved local shard paths self._local_shard_paths = {} + # Cache for open safetensors file readers + self._open_files = {} # Use a lock to serialize heavy RAM operations, but NOT downloads self._ram_lock = threading.Lock() self._initialize_index() def __getstate__(self): - """Allows pickling/copying by excluding the non-pickleable lock.""" + """Allows pickling/copying by excluding the non-pickleable lock and files.""" state = self.__dict__.copy() del state["_ram_lock"] + if "_open_files" in state: + del state["_open_files"] return state def __setstate__(self, state): - """Restores state after pickling/copying and recreates a new lock.""" + """Restores state after pickling/copying and recreates a new lock and file cache.""" self.__dict__.update(state) self._ram_lock = threading.Lock() + self._open_files = {} def _initialize_index(self): """Fetches and parses the Hugging Face model index file to build a shard map.""" @@ -206,7 +212,17 @@ def get_tensor(self, key: str) -> np.ndarray: # This prevents multiple threads from simultaneously allocating large chunks of RAM. with self._ram_lock: with safe_open(local_path, framework="np", device="cpu") as f: - return f.get_tensor(key) + tensor = f.get_tensor(key) + # Prune older downloaded shards if downloading remotely to prevent disk/RAM exhaustion + if not self.is_local and len(self._local_shard_paths) > 3: + oldest_shard = next(iter(self._local_shard_paths)) + oldest_path = self._local_shard_paths.pop(oldest_shard) + if os.path.exists(oldest_path): + try: + os.remove(oldest_path) + except OSError: + pass + return tensor class LazyTensor: @@ -448,7 +464,28 @@ def _build_single_axis_stacked_tensor( if isinstance(hf_key_single, (list, tuple)): hf_tensor_numpy = tuple(tensor_getter_fn(k) for k in hf_key_single) else: - hf_tensor_numpy = tensor_getter_fn(hf_key_single) + try: + hf_tensor_numpy = tensor_getter_fn(hf_key_single) + except (ValueError, KeyError) as e: + if "indexer" in str(hf_key_single) and str(hf_key_single).startswith("model.layers."): + m = re.match(r"model\.layers\.(\d+)\.(.+)", str(hf_key_single)) + if m: + layer_idx = int(m.group(1)) + rest = m.group(2) + # Search backwards for the closest preceding donor layer containing the key + for candidate_idx in range(layer_idx - 1, -1, -1): + donor_key = f"model.layers.{candidate_idx}.{rest}" + try: + hf_tensor_numpy = tensor_getter_fn(donor_key) + break + except Exception: # pylint: disable=broad-exception-caught + continue + else: + hf_tensor_numpy = np.zeros(mt_slice_shape, dtype=np.float32) + else: + raise e + else: + raise e processed_hf_tensor = apply_hook_fns(hf_tensor_numpy, mt_slice_shape, hook_fns) tensors_to_stack.append(processed_hf_tensor) @@ -984,6 +1021,16 @@ def main( def _eager_getter(key): if key not in hf_state_dict_numpy: + if getattr(config, "use_index_share", False) and "indexer" in key and key.startswith("model.layers."): + m = re.match(r"model\.layers\.(\d+)\.(.+)", key) + if m: + layer_idx = int(m.group(1)) + rest = m.group(2) + # Search backwards for the closest preceding donor layer containing the key + for candidate_idx in range(layer_idx - 1, -1, -1): + donor_key = f"model.layers.{candidate_idx}.{rest}" + if donor_key in hf_state_dict_numpy: + return _eager_getter(donor_key) raise ValueError(f"HuggingFace key {key} not found in state_dict.") v = hf_state_dict_numpy[key] # target dtype is "float32" @@ -1002,6 +1049,29 @@ def _eager_getter(key): tensor_getter = _eager_getter + if getattr(config, "use_index_share", False): + orig_tensor_getter = tensor_getter + + def _index_share_tensor_getter(key): + try: + return orig_tensor_getter(key) + except (ValueError, KeyError) as e: + if "indexer" in key and key.startswith("model.layers."): + m = re.match(r"model\.layers\.(\d+)\.(.+)", key) + if m: + layer_idx = int(m.group(1)) + rest = m.group(2) + # Search backwards for the closest preceding donor layer containing the key + for candidate_idx in range(layer_idx - 1, -1, -1): + donor_key = f"model.layers.{candidate_idx}.{rest}" + try: + return orig_tensor_getter(donor_key) + except (ValueError, KeyError): + continue + raise e + + tensor_getter = _index_share_tensor_getter + if is_merge_mode: tensor_getter = _setup_merge_mode_getter(tensor_getter, config, hf_lora_adapter_path, revision) diff --git a/src/maxtext/checkpoint_conversion/utils/hf_model_configs.py b/src/maxtext/checkpoint_conversion/utils/hf_model_configs.py index 89abd56d4c..d493efb67f 100644 --- a/src/maxtext/checkpoint_conversion/utils/hf_model_configs.py +++ b/src/maxtext/checkpoint_conversion/utils/hf_model_configs.py @@ -1884,8 +1884,27 @@ def __init__(self, **kwargs): qwen3_vl_30b_a3b_config = PTConfig(**qwen3_vl_30b_a3b_dict) +glm5_1_744b_dict = { + "architectures": ["GlmMoeDsaForCausalLM"], + "num_hidden_layers": 78, + "first_k_dense_replace": 3, + "n_routed_experts": 256, +} +glm5_1_744b_config = transformers.DeepseekV3Config(**glm5_1_744b_dict) + +glm5_2_744b_dict = { + "architectures": ["GlmMoeDsaForCausalLM"], + "num_hidden_layers": 78, + "first_k_dense_replace": 3, + "n_routed_experts": 256, +} +glm5_2_744b_config = transformers.DeepseekV3Config(**glm5_2_744b_dict) + + # {maxtext model name: hf model config} HF_MODEL_CONFIGS = { + "glm5.1-744b": glm5_1_744b_config, + "glm5.2-744b": glm5_2_744b_config, "gemma2-2b": gemma2_2b_config, "gemma2-9b": gemma2_9b_config, "gemma2-27b": gemma2_27b_config, diff --git a/src/maxtext/checkpoint_conversion/utils/hf_shape.py b/src/maxtext/checkpoint_conversion/utils/hf_shape.py index 85dd1d6ea0..9b2f7094a6 100644 --- a/src/maxtext/checkpoint_conversion/utils/hf_shape.py +++ b/src/maxtext/checkpoint_conversion/utils/hf_shape.py @@ -1312,6 +1312,8 @@ def DEEPSEEKV4_HF_WEIGHTS_TO_SHAPE(config): "deepseek3-671b": DEEPSEEK_HF_WEIGHTS_TO_SHAPE, "deepseek3.2-671b": DEEPSEEK_HF_WEIGHTS_TO_SHAPE, "deepseek4-284b": DEEPSEEKV4_HF_WEIGHTS_TO_SHAPE, + "glm5.1-744b": DEEPSEEK_HF_WEIGHTS_TO_SHAPE, + "glm5.2-744b": DEEPSEEK_HF_WEIGHTS_TO_SHAPE, "gpt-oss-20b": GPT_OSS_HF_WEIGHTS_TO_SHAPE, "gpt-oss-120b": GPT_OSS_HF_WEIGHTS_TO_SHAPE, "mixtral-8x7b": MIXTRAL_HF_WEIGHTS_TO_SHAPE, diff --git a/src/maxtext/checkpoint_conversion/utils/param_mapping.py b/src/maxtext/checkpoint_conversion/utils/param_mapping.py index 94e96173f1..2ae6b57869 100644 --- a/src/maxtext/checkpoint_conversion/utils/param_mapping.py +++ b/src/maxtext/checkpoint_conversion/utils/param_mapping.py @@ -1795,6 +1795,319 @@ def DEEPSEEK_NNX_TO_VLLM_PARAM_HOOK_FN(): return {} +def GLM_MAXTEXT_TO_HF_PARAM_MAPPING(config, maxtext_config, scan_layers=False): + """Generates mapping from MaxText GLM-5.1/5.2 to Hugging Face weight paths.""" + num_main_layers = config["num_hidden_layers"] + first_num_dense_layers = config["first_k_dense_replace"] + num_experts = config.get("n_routed_experts", 0) + + # Mapping for non-layer-specific weights + mapping = { + "params-token_embedder-embedding": "model.embed_tokens.weight", + "params-decoder-decoder_norm-scale": "model.norm.weight", + "params-decoder-logits_dense-kernel": "lm_head.weight", + } + # Attention keys are shared by both dense and MoE + attention_keys = { + "pre_self_attention_layer_norm-scale": "input_layernorm.weight", + "post_self_attention_layer_norm-scale": "post_attention_layernorm.weight", + "self_attention-kv_norm-scale": "self_attn.kv_a_layernorm.weight", + "self_attention-wkv_a-kernel": "self_attn.kv_a_proj_with_mqa.weight", + "self_attention-wkv_b-kernel": "self_attn.kv_b_proj.weight", + "self_attention-out-kernel": "self_attn.o_proj.weight", + "self_attention-q_norm-scale": "self_attn.q_a_layernorm.weight", + "self_attention-wq_a-kernel": "self_attn.q_a_proj.weight", + "self_attention-wq_b-kernel": "self_attn.q_b_proj.weight", + "self_attention-indexer-k_norm-bias": "self_attn.indexer.k_norm.bias", + "self_attention-indexer-k_norm-scale": "self_attn.indexer.k_norm.weight", + "self_attention-indexer-weights_proj-kernel": "self_attn.indexer.weights_proj.weight", + "self_attention-indexer-wk-kernel": "self_attn.indexer.wk.weight", + "self_attention-indexer-wq_b-kernel": "self_attn.indexer.wq_b.weight", + } + # Dense Layers + dense_layer_keys = attention_keys | { + "mlp-wi_0-kernel": "mlp.gate_proj.weight", + "mlp-wi_1-kernel": "mlp.up_proj.weight", + "mlp-wo-kernel": "mlp.down_proj.weight", + } + # MoE Layers + moe_layer_keys = attention_keys | { + "DeepSeekMoeBlock_0-shared_experts-wi_0-kernel": "mlp.shared_experts.gate_proj.weight", + "DeepSeekMoeBlock_0-shared_experts-wi_1-kernel": "mlp.shared_experts.up_proj.weight", + "DeepSeekMoeBlock_0-shared_experts-wo-kernel": "mlp.shared_experts.down_proj.weight", + "DeepSeekMoeBlock_0-MoeBlock_0-gate-kernel": "mlp.gate.weight", + "DeepSeekMoeBlock_0-MoeBlock_0-gate-bias": "mlp.gate.e_score_correction_bias", + } + # MoE Experts (nested list mapping: [[e0_l0, e0_l1..], [e1_l0, e1_l1..]..]) + moe_expert_keys = { + "DeepSeekMoeBlock_0-MoeBlock_0-wi_0": "gate_proj.weight", + "DeepSeekMoeBlock_0-MoeBlock_0-wi_1": "up_proj.weight", + "DeepSeekMoeBlock_0-MoeBlock_0-wo": "down_proj.weight", + } + + # scan + if scan_layers: + for maxtext_key, hf_key in dense_layer_keys.items(): + mapping[f"params-decoder-dense_layers-{maxtext_key}"] = [ + f"model.layers.{i}.{hf_key}" for i in range(first_num_dense_layers) + ] + + for maxtext_key, hf_key in moe_layer_keys.items(): + mapping[f"params-decoder-moe_layers-{maxtext_key}"] = [ + f"model.layers.{i}.{hf_key}" for i in range(first_num_dense_layers, num_main_layers) + ] + + for maxtext_key, hf_key in moe_expert_keys.items(): + mapping[f"params-decoder-moe_layers-{maxtext_key}"] = [ + [f"model.layers.{i}.mlp.experts.{e}.{hf_key}" for i in range(first_num_dense_layers, num_main_layers)] + for e in range(num_experts) + ] + # unscan + else: + for i in range(first_num_dense_layers): + for maxtext_key, hf_key in dense_layer_keys.items(): + mapping[f"params-decoder-dense_layers_{i}-{maxtext_key}"] = f"model.layers.{i}.{hf_key}" + + for i in range(first_num_dense_layers, num_main_layers): + moe_layer_idx = i - first_num_dense_layers + + for maxtext_key, hf_key in moe_layer_keys.items(): + mapping[f"params-decoder-moe_layers_{moe_layer_idx}-{maxtext_key}"] = f"model.layers.{i}.{hf_key}" + + for maxtext_key, hf_key in moe_expert_keys.items(): + mapping[f"params-decoder-moe_layers_{moe_layer_idx}-{maxtext_key}"] = [ + f"model.layers.{i}.mlp.experts.{e}.{hf_key}" for e in range(num_experts) + ] + return mapping + + +def GLM_MAXTEXT_TO_HF_PARAM_HOOK_FN(config, maxtext_config, scan_layers=False, saving_to_hf=False): + """Creates parameter transformation functions for GLM-5.1 & GLM-5.2.""" + + def reshape_kernel(input_tensor, target_shape): + """Reshapes and transposes standard 2D, 3D, and 4D linear kernels.""" + if saving_to_hf: + if input_tensor.ndim == 4: + return input_tensor.transpose(0, 1, 3, 2) + elif input_tensor.ndim == 3: + return input_tensor.transpose(1, 2, 0) + elif input_tensor.ndim == 2: + return input_tensor.T + return input_tensor + else: + if input_tensor.ndim == 4: + return input_tensor.transpose(0, 1, 3, 2) + elif input_tensor.ndim == 3: + return input_tensor.transpose(2, 0, 1) + elif input_tensor.ndim == 2: + return input_tensor.T + return input_tensor + + def reshape_wkv_b_kernel(input_tensor, target_shape): + """Reshapes and transposes wkv_b kernel weights between MaxText and HF GLM-5.2. + + HF kv_b_proj.weight shape is [num_heads * (qk_nope_head_dim + v_head_dim), kv_lora_rank]. + In HF GlmMoeDsa, the weight is arranged per-head: [num_heads, qk_nope_head_dim + v_head_dim, kv_lora_rank]. + JAX expects wkv_b shape [kv_lora_rank, num_heads, qk_nope_head_dim + v_head_dim] + or scanned [num_layers, kv_lora_rank, num_heads, qk_nope_head_dim + v_head_dim]. + """ + num_heads = maxtext_config.num_query_heads + qk_nope_head_dim = maxtext_config.qk_nope_head_dim + v_head_dim = maxtext_config.v_head_dim + head_dim = qk_nope_head_dim + v_head_dim + + if saving_to_hf: + # JAX -> HF + if input_tensor.ndim == 4: # [L, kv_lora_rank, num_heads, head_dim] + return input_tensor.transpose(0, 2, 3, 1).reshape( + input_tensor.shape[0], num_heads * head_dim, input_tensor.shape[1] + ) + return input_tensor.transpose(1, 2, 0).reshape(num_heads * head_dim, input_tensor.shape[0]) + else: + # HF -> JAX + if input_tensor.ndim == 3: # [L, num_heads * head_dim, kv_lora_rank] + return input_tensor.reshape(input_tensor.shape[0], num_heads, head_dim, input_tensor.shape[-1]).transpose( + 0, 3, 1, 2 + ) + return input_tensor.reshape(num_heads, head_dim, input_tensor.shape[-1]).transpose(2, 0, 1) + + def reshape_wq_b_kernel(input_tensor, target_shape): + """Reshapes and transposes wq_b kernel weights between MaxText and HF GLM-5.2. + + HF q_b_proj.weight shape is [num_heads * (qk_nope_head_dim + qk_rope_head_dim), q_lora_rank]. + In HF GlmMoeDsa, the weight is arranged per-head: [num_heads, qk_nope_head_dim + qk_rope_head_dim, q_lora_rank]. + JAX expects wq_b shape [q_lora_rank, num_heads, qk_nope_head_dim + qk_rope_head_dim] + or scanned [num_layers, q_lora_rank, num_heads, qk_nope_head_dim + qk_rope_head_dim]. + """ + num_heads = maxtext_config.num_query_heads + qk_nope_head_dim = maxtext_config.qk_nope_head_dim + qk_rope_head_dim = maxtext_config.qk_rope_head_dim + head_dim = qk_nope_head_dim + qk_rope_head_dim + + if saving_to_hf: + # JAX -> HF + if input_tensor.ndim == 4: # [L, q_lora_rank, num_heads, head_dim] + return input_tensor.transpose(0, 2, 3, 1).reshape( + input_tensor.shape[0], num_heads * head_dim, input_tensor.shape[1] + ) + return input_tensor.transpose(1, 2, 0).reshape(num_heads * head_dim, input_tensor.shape[0]) + else: + # HF -> JAX + if input_tensor.ndim == 3: # [L, num_heads * head_dim, q_lora_rank] + return input_tensor.reshape(input_tensor.shape[0], num_heads, head_dim, input_tensor.shape[-1]).transpose( + 0, 3, 1, 2 + ) + return input_tensor.reshape(num_heads, head_dim, input_tensor.shape[-1]).transpose(2, 0, 1) + + def reshape_indexer_wq_b_kernel(input_tensor, target_shape): + """Reshapes and transposes indexer wq_b kernel weights. + + HF indexer.wq_b.weight has shape [4096, 2048]. + JAX NNX scanned indexer-wq_b-kernel expects [num_layers=78, in_dim=2048, num_heads=32, head_dim=128]. + JAX Linen scanned indexer-wq_b-kernel expects [2048, num_layers=78, 32, 128]. + """ + num_heads = maxtext_config.indexer_n_heads + head_dim = maxtext_config.indexer_head_dim + + if saving_to_hf: + # JAX -> HF + if input_tensor.ndim == 4: + if input_tensor.shape[0] == num_heads or input_tensor.shape[0] == input_tensor.shape[-1]: + transposed = input_tensor.transpose(1, 2, 3, 0) + else: + transposed = input_tensor.transpose(0, 2, 3, 1) + reshaped = transposed.reshape(transposed.shape[0], num_heads * head_dim, transposed.shape[-1]) + if len(target_shape) == 2: + return reshaped[0].T + return reshaped.transpose(0, 2, 1) + else: + transposed = input_tensor.transpose(1, 2, 0) + return transposed.reshape(num_heads * head_dim, transposed.shape[-1]).T + else: + # HF -> JAX + if input_tensor.ndim == 2: + input_tensor = input_tensor[None, :, :] + # Reshape [L, 4096, 2048] -> [L, H, D, I] + reshaped = input_tensor.reshape(input_tensor.shape[0], num_heads, head_dim, input_tensor.shape[-1]) + if len(target_shape) == 4: + if target_shape[0] == input_tensor.shape[-1]: # [I, L, H, D] (Linen) + return reshaped.transpose(3, 0, 1, 2) + return reshaped.transpose(0, 3, 1, 2) # [L, I, H, D] (NNX) + elif len(target_shape) == 3: + if target_shape[0] == input_tensor.shape[-1]: + return reshaped.transpose(3, 0, 1, 2)[:, 0, :, :] + return reshaped.transpose(0, 3, 1, 2)[0, :, :, :] + return reshaped + + def reshape_out_kernel(input_tensor, target_shape): + """Reshapes and transposes out kernel weights. + + HF o_proj.weight has shape [6144, 16384]. + JAX NNX scanned out-kernel expects [num_layers=78, num_heads=64, head_dim=256, hidden_dim=6144]. + JAX Linen scanned out-kernel expects [num_heads=64, num_layers=78, head_dim=256, hidden_dim=6144]. + """ + num_heads = maxtext_config.num_query_heads + v_head_dim = maxtext_config.v_head_dim + + if saving_to_hf: + # JAX -> HF + if input_tensor.ndim == 4: + if input_tensor.shape[0] == num_heads: + transposed = input_tensor.transpose(1, 3, 0, 2) + else: + transposed = input_tensor.transpose(0, 3, 1, 2) + reshaped = transposed.reshape(transposed.shape[0], transposed.shape[1], num_heads * v_head_dim) + if len(target_shape) == 2: + return reshaped[0].T + return reshaped.transpose(0, 2, 1) + else: + transposed = input_tensor.transpose(2, 0, 1) + return transposed.reshape(transposed.shape[0], num_heads * v_head_dim).T + else: + # HF -> JAX + if input_tensor.ndim == 2: + input_tensor = input_tensor[None, :, :] + # Reshape [L, 6144, 16384] -> [L, I, H, D] + reshaped = input_tensor.reshape(input_tensor.shape[0], input_tensor.shape[1], num_heads, v_head_dim) + if len(target_shape) == 4: + if target_shape[0] == num_heads: + return reshaped.transpose(2, 0, 3, 1) # [H, L, D, I] (Linen) + return reshaped.transpose(0, 2, 3, 1) # [L, H, D, I] (NNX) + elif len(target_shape) == 3: + if target_shape[0] == num_heads: + return reshaped.transpose(2, 0, 3, 1)[:, 0, :, :] + return reshaped.transpose(0, 2, 3, 1)[0, :, :, :] + return reshaped + + num_main_layers = config["num_hidden_layers"] + first_num_dense_layers = config["first_k_dense_replace"] + + mapping = { + "params-decoder-logits_dense-kernel": reshape_kernel, + } + + attention_need_reshape = { + "self_attention-wkv_a-kernel", # transpose + "self_attention-query-kernel", + "self_attention-wq_a-kernel", # transpose + "self_attention-indexer-weights_proj-kernel", # transpose + "self_attention-indexer-wk-kernel", # transpose + } + + dense_need_reshape = attention_need_reshape | { + "mlp-wi_0-kernel", # transpose + "mlp-wi_1-kernel", # transpose + "mlp-wo-kernel", # transpose + } + + moe_need_reshape = attention_need_reshape | { + "DeepSeekMoeBlock_0-shared_experts-wi_0-kernel", # transpose + "DeepSeekMoeBlock_0-shared_experts-wi_1-kernel", # transpose + "DeepSeekMoeBlock_0-shared_experts-wo-kernel", # transpose + "DeepSeekMoeBlock_0-MoeBlock_0-gate-kernel", # transpose + "DeepSeekMoeBlock_0-MoeBlock_0-gate-bias", # transpose + "DeepSeekMoeBlock_0-MoeBlock_0-wi_0", # transpose + "DeepSeekMoeBlock_0-MoeBlock_0-wi_1", # transpose + "DeepSeekMoeBlock_0-MoeBlock_0-wo", # transpose + } + + # scan + if scan_layers: + for key in dense_need_reshape: + mapping[f"params-decoder-dense_layers-{key}"] = reshape_kernel + for key in moe_need_reshape: + mapping[f"params-decoder-moe_layers-{key}"] = reshape_kernel + mapping["params-decoder-dense_layers-self_attention-wkv_b-kernel"] = reshape_wkv_b_kernel + mapping["params-decoder-moe_layers-self_attention-wkv_b-kernel"] = reshape_wkv_b_kernel + mapping["params-decoder-dense_layers-self_attention-wq_b-kernel"] = reshape_wq_b_kernel + mapping["params-decoder-moe_layers-self_attention-wq_b-kernel"] = reshape_wq_b_kernel + mapping["params-decoder-dense_layers-self_attention-indexer-wq_b-kernel"] = reshape_indexer_wq_b_kernel + mapping["params-decoder-moe_layers-self_attention-indexer-wq_b-kernel"] = reshape_indexer_wq_b_kernel + mapping["params-decoder-dense_layers-self_attention-out-kernel"] = reshape_out_kernel + mapping["params-decoder-moe_layers-self_attention-out-kernel"] = reshape_out_kernel + # unscan + else: + for i in range(first_num_dense_layers): + for key in dense_need_reshape: + mapping[f"params-decoder-dense_layers_{i}-{key}"] = reshape_kernel + mapping[f"params-decoder-dense_layers_{i}-self_attention-wkv_b-kernel"] = reshape_wkv_b_kernel + mapping[f"params-decoder-dense_layers_{i}-self_attention-wq_b-kernel"] = reshape_wq_b_kernel + mapping[f"params-decoder-dense_layers_{i}-self_attention-indexer-wq_b-kernel"] = reshape_indexer_wq_b_kernel + mapping[f"params-decoder-dense_layers_{i}-self_attention-out-kernel"] = reshape_out_kernel + for i in range(first_num_dense_layers, num_main_layers): + moe_layer_idx = i - first_num_dense_layers + for key in moe_need_reshape: + mapping[f"params-decoder-moe_layers_{moe_layer_idx}-{key}"] = reshape_kernel + mapping[f"params-decoder-moe_layers_{moe_layer_idx}-self_attention-wkv_b-kernel"] = reshape_wkv_b_kernel + mapping[f"params-decoder-moe_layers_{moe_layer_idx}-self_attention-wq_b-kernel"] = reshape_wq_b_kernel + mapping[f"params-decoder-moe_layers_{moe_layer_idx}-self_attention-indexer-wq_b-kernel"] = ( + reshape_indexer_wq_b_kernel + ) + mapping[f"params-decoder-moe_layers_{moe_layer_idx}-self_attention-out-kernel"] = reshape_out_kernel + + return mapping + + def GPT_OSS_MAXTEXT_TO_HF_PARAM_MAPPING(config, maxtext_config, scan_layers=False): """Generates mapping from MaxText gpt-oss to Hugging Face weight paths. @@ -4257,6 +4570,8 @@ def mhc_concat_scale(input_tensors, target_shape=None): "deepseek3-671b": DEEPSEEK_MAXTEXT_TO_HF_PARAM_MAPPING, "deepseek3.2-671b": DEEPSEEK_MAXTEXT_TO_HF_PARAM_MAPPING, "deepseek4-284b": DEEPSEEKV4_MAXTEXT_TO_HF_PARAM_MAPPING, + "glm5.1-744b": GLM_MAXTEXT_TO_HF_PARAM_MAPPING, + "glm5.2-744b": GLM_MAXTEXT_TO_HF_PARAM_MAPPING, "gpt-oss-20b": GPT_OSS_MAXTEXT_TO_HF_PARAM_MAPPING, "gpt-oss-120b": GPT_OSS_MAXTEXT_TO_HF_PARAM_MAPPING, "qwen3-omni-30b-a3b": QWEN3_OMNI_MOE_MAXTEXT_TO_HF_PARAM_MAPPING, @@ -4312,6 +4627,8 @@ def mhc_concat_scale(input_tensors, target_shape=None): "deepseek3.2-671b": DEEPSEEK_MAXTEXT_TO_HF_PARAM_HOOK_FN, "deepseek4-tiny": DEEPSEEKV4_MAXTEXT_TO_HF_PARAM_HOOK_FN, "deepseek4-284b": DEEPSEEKV4_MAXTEXT_TO_HF_PARAM_HOOK_FN, + "glm5.1-744b": GLM_MAXTEXT_TO_HF_PARAM_HOOK_FN, + "glm5.2-744b": GLM_MAXTEXT_TO_HF_PARAM_HOOK_FN, "gpt-oss-20b": GPT_OSS_TO_HF_PARAM_HOOK_FN, "gpt-oss-120b": GPT_OSS_TO_HF_PARAM_HOOK_FN, "qwen3-omni-30b-a3b": QWEN3_OMNI_MOE_MAXTEXT_TO_HF_PARAM_HOOK_FN, diff --git a/src/maxtext/common/common_types.py b/src/maxtext/common/common_types.py index 77f93a63d7..41177aea8a 100644 --- a/src/maxtext/common/common_types.py +++ b/src/maxtext/common/common_types.py @@ -116,6 +116,7 @@ class DecoderBlockType(enum.Enum): OLMO3 = "olmo3" DEEPSEEK4 = "deepseek4" ENVY = "envy" + GLM5 = "glm5" class VisionEncoderBlockType(enum.Enum): diff --git a/src/maxtext/common/gcloud_stub.py b/src/maxtext/common/gcloud_stub.py index c87ba4123c..bdd48a0ab8 100644 --- a/src/maxtext/common/gcloud_stub.py +++ b/src/maxtext/common/gcloud_stub.py @@ -105,22 +105,178 @@ def __init__( self.log_prob = log_prob self.samples_per_slot = samples_per_slot - # Tokenizer placeholders (unused in decoupled tests due to runtime guard). - class TokenizerParameters: # pragma: no cover - placeholder + def get_result_at_slot(self, slot: int): + """Extracts the result tokens at a particular slot.""" + if self.data is not None and self.tokens_idx is not None: + if isinstance(self.tokens_idx, tuple) and len(self.tokens_idx) == 2: + tokens = self.data[slot, self.tokens_idx[0] : self.tokens_idx[1]] + else: + tokens = self.data[slot, self.tokens_idx] + else: + tokens = self.data + + if self.data is not None and self.valid_idx is not None: + if isinstance(self.valid_idx, tuple) and len(self.valid_idx) == 2: + valid = self.data[slot, self.valid_idx[0] : self.valid_idx[1]] + else: + valid = self.data[slot, self.valid_idx] + else: + valid = None + + if self.data is not None and self.length_idx is not None: + if isinstance(self.length_idx, tuple) and len(self.length_idx) == 2: + length = self.data[slot, self.length_idx[0] : self.length_idx[1]] + else: + length = self.data[slot, self.length_idx] + else: + length = None + + log_prob = self.log_prob[slot] if self.log_prob is not None else None + return SimpleNamespace( + tokens=tokens, + valid=valid, + length=length, + log_prob=log_prob, + ) + + def _tree_flatten(self): + children = (self.data, self.log_prob) + aux_data = (self.tokens_idx, self.valid_idx, self.length_idx, self.samples_per_slot) + return children, aux_data + + @classmethod + def _tree_unflatten(cls, aux_data, children): + data, log_prob = children + tokens_idx, valid_idx, length_idx, samples_per_slot = aux_data + return cls( + data=data, + log_prob=log_prob, + tokens_idx=tokens_idx, + valid_idx=valid_idx, + length_idx=length_idx, + samples_per_slot=samples_per_slot, + ) - def __init__(self, *a, **k): - pass + try: + import jax # pylint: disable=import-outside-toplevel + + jax.tree_util.register_pytree_node( + ResultTokens, + ResultTokens._tree_flatten, # pylint: disable=protected-access + ResultTokens._tree_unflatten, # pylint: disable=protected-access + ) + except Exception: # pylint: disable=broad-exception-caught + pass - class TokenizerType: # emulate enum descriptor access pattern - DESCRIPTOR = SimpleNamespace(values_by_name={}) + class TokenizerParameters: + """Container for tokenizer parameters.""" + + def __init__(self, path=None, tokenizer_type=None, access_token=None, use_chat_template=False, extra_ids=0, **kwargs): + del kwargs + self.path = path + self.tokenizer_type = tokenizer_type + self.access_token = access_token + self.use_chat_template = use_chat_template + self.extra_ids = extra_ids + + class TokenizerType: + """Enum emulator for tokenizer types.""" + + tiktoken = 1 + sentencepiece = 2 + huggingface = 3 + DESCRIPTOR = SimpleNamespace( + values_by_name={ + "tiktoken": SimpleNamespace(number=1), + "sentencepiece": SimpleNamespace(number=2), + "huggingface": SimpleNamespace(number=3), + } + ) + + class HuggingFaceTokenizer: + """Fallback HuggingFace tokenizer when JetStream is not installed.""" + + def __init__(self, metadata): + import transformers # pylint: disable=import-outside-toplevel + + try: + self.tokenizer = transformers.AutoTokenizer.from_pretrained( + metadata.path, + token=metadata.access_token or None, + trust_remote_code=True, + ) + except Exception: # pylint: disable=broad-exception-caught + try: + self.tokenizer = transformers.PreTrainedTokenizerFast.from_pretrained( + metadata.path, + token=metadata.access_token or None, + ) + except Exception: # pylint: disable=broad-exception-caught + try: + import huggingface_hub # pylint: disable=import-outside-toplevel + + tok_file = metadata.path + if not os.path.exists(tok_file): + tok_file = huggingface_hub.hf_hub_download( + repo_id=metadata.path, + filename="tokenizer.json", + token=metadata.access_token or None, + ) + self.tokenizer = transformers.PreTrainedTokenizerFast(tokenizer_file=tok_file) + except Exception: # pylint: disable=broad-exception-caught + self.tokenizer = transformers.AutoTokenizer.from_pretrained( + "THUDM/glm-4-9b-chat", + token=metadata.access_token or None, + trust_remote_code=True, + ) + + self.pad_token_id = getattr(self.tokenizer, "pad_token_id", None) + self.eos_token_id = getattr(self.tokenizer, "eos_token_id", None) + self.bos_token_id = getattr(self.tokenizer, "bos_token_id", None) + if self.pad_token_id is None and hasattr(self.tokenizer, "token_to_id"): + self.eos_token_id = self.tokenizer.token_to_id("<|endoftext|>") + self.pad_token_id = self.eos_token_id + try: + self.tokenizer.pad_token_id = self.pad_token_id + self.tokenizer.eos_token_id = self.eos_token_id + except Exception: # pylint: disable=broad-exception-caught + pass - config_lib = SimpleNamespace() # not used directly in decoupled tests + def encode(self, text, is_bos=True, prefill_lengths=None): + """Encodes text to token IDs with optional padding.""" + import numpy as np # pylint: disable=import-outside-toplevel + + if hasattr(self.tokenizer, "encode"): + res = self.tokenizer.encode(text) + token_ids = res.ids if hasattr(res, "ids") else res + else: + token_ids = [] + bos_id = getattr(self.tokenizer, "bos_token_id", None) + if is_bos and bos_id is not None: + token_ids = [bos_id] + list(token_ids) + true_length = len(token_ids) + target_len = prefill_lengths[0] if prefill_lengths else true_length + pad_id = getattr(self.tokenizer, "pad_token_id", 0) or 0 + padded = list(token_ids) + [pad_id] * max(0, target_len - true_length) + return np.array(padded[:target_len], dtype=np.int32), true_length + + def decode(self, token_ids): + """Decodes token IDs back to string.""" + if hasattr(token_ids, "tolist"): + token_ids = token_ids.tolist() + if hasattr(self.tokenizer, "decode"): + try: + return self.tokenizer.decode(token_ids, skip_special_tokens=True) + except TypeError: + return self.tokenizer.decode(token_ids) + return "" + + config_lib = SimpleNamespace() engine_api = SimpleNamespace(Engine=Engine, ResultTokens=ResultTokens) - token_utils = SimpleNamespace() # build_tokenizer guarded in MaxEngine when decoupled - tokenizer_api = SimpleNamespace() # placeholder + token_utils = SimpleNamespace(HuggingFaceTokenizer=HuggingFaceTokenizer) + tokenizer_api = SimpleNamespace() token_params_ns = SimpleNamespace(TokenizerParameters=TokenizerParameters, TokenizerType=TokenizerType) - # Mark these stub namespaces so callers can detect stubbed jetstream components. setattr(config_lib, "_IS_STUB", True) setattr(engine_api, "_IS_STUB", True) setattr(token_utils, "_IS_STUB", True) diff --git a/src/maxtext/common/profiler.py b/src/maxtext/common/profiler.py index f6f56c372e..e633ab2b4a 100644 --- a/src/maxtext/common/profiler.py +++ b/src/maxtext/common/profiler.py @@ -75,6 +75,13 @@ def __init__(self, config, offset_step=0): if advanced_config: self.profiling_options.advanced_configuration = advanced_config + def is_active(self, step=None): + if self.mode == "": + return False + if step is not None: + return self.start_initial_profile_step <= step <= self.finished_initial_profile_step + return getattr(self, "_active", False) + def maybe_activate_profiler(self, step, state): """Conditionally activates the profiler based on the current step. This method checks if the current training step matches the step designated @@ -83,13 +90,14 @@ def maybe_activate_profiler(self, step, state): """ if self.mode != "" and (step == self.start_initial_profile_step or self.should_activate_periodic_profile(step)): optional_postfix = f"step_{step}" if self.profile_period > 0 else "" + self._active = True self.activate(blocking_object=state, optional_postfix=optional_postfix) def activate(self, blocking_object=None, optional_postfix=""): """Start the profiler. nsys profiler becomes no-op when libcudart.so is not available on the system.""" if self.profile_cleanly and blocking_object is not None: - jax.block_until_ready(blocking_object) + jax.tree_util.tree_map(lambda x: x.block_until_ready() if hasattr(x, "block_until_ready") else x, blocking_object) if self.managed_mldiagnostics and self.mode == "xplane": # Handle the special profiling logic for managed_mldiagnostics @@ -121,13 +129,14 @@ def maybe_deactivate_profiler(self, step, state): deactivating a periodic profile. """ if self.mode != "" and (step == self.finished_initial_profile_step or self.should_deactivate_periodic_profile(step)): + self._active = False self.deactivate(blocking_object=state) def deactivate(self, blocking_object=None): """End the profiler. The result is uploaded to the output bucket.""" if self.profile_cleanly and blocking_object is not None: - jax.block_until_ready(blocking_object) + jax.tree_util.tree_map(lambda x: x.block_until_ready() if hasattr(x, "block_until_ready") else x, blocking_object) if self.managed_mldiagnostics and self.mode == "xplane": # Handle the special profileing logic for managed_mldiagnostics diff --git a/src/maxtext/configs/base.yml b/src/maxtext/configs/base.yml index d8efa17359..62106f1dbe 100644 --- a/src/maxtext/configs/base.yml +++ b/src/maxtext/configs/base.yml @@ -444,6 +444,11 @@ indexer_sparse_training: false # Multiplier for the indexer KL divergence loss indexer_loss_scaling_factor: 0.0 +# GLM-5.2 Cross-Layer IndexCache / IndexShare +use_index_share: false +index_share_pattern: "FSSS" +prune_shared_indexers: true + # MLA parameters q_lora_rank: 0 kv_lora_rank: 512 diff --git a/src/maxtext/configs/models/glm5.1-744b.yml b/src/maxtext/configs/models/glm5.1-744b.yml new file mode 100644 index 0000000000..5f23c59761 --- /dev/null +++ b/src/maxtext/configs/models/glm5.1-744b.yml @@ -0,0 +1,68 @@ +# 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. + +# model config for GLM-5.1 - 744B (Mixture of Experts) + +base_emb_dim: 6144 +base_num_query_heads: 64 +base_num_kv_heads: 64 +base_mlp_dim: 12288 +base_moe_mlp_dim: 2048 + +base_num_decoder_layers: 78 + +first_num_dense_layers: 3 +mlp_activations: ["silu","linear"] +vocab_size: 154880 +enable_dropout: false +logits_via_embedding: false +normalization_layer_epsilon: 1.0e-5 +num_experts: 256 +num_experts_per_tok: 8 +shared_experts: 1 +routed_scaling_factor: 2.5 +routed_score_func: "sigmoid" +routed_bias: true +norm_topk_prob: true +decoder_block: "glm5" +dtype: "bfloat16" +weight_dtype: "bfloat16" + +# Multi-head Latent Attention (MLA) +attention_type: "mla" +attention: "dot_product" +q_lora_rank: 2048 +kv_lora_rank: 512 +qk_nope_head_dim: 192 +qk_rope_head_dim: 64 +v_head_dim: 256 + +# RoPE +mscale: 1.0 +rope_type: "yarn" +rope_max_timescale: 1000000 # "rope_theta": 1000000 +max_position_embeddings: 202752 +original_max_position_embeddings: 4096 +rope_factor: 1.0 +beta_fast: 32 +beta_slow: 1 +rope_interleave: true +rope_truncate: true +rope_attention_scaling: false + +# Indexer for Dynamic Sparse Attention (DSA) +use_indexer: true +indexer_n_heads: 32 +indexer_head_dim: 128 +indexer_topk: 2048 diff --git a/src/maxtext/configs/models/glm5.2-744b.yml b/src/maxtext/configs/models/glm5.2-744b.yml new file mode 100644 index 0000000000..e520015667 --- /dev/null +++ b/src/maxtext/configs/models/glm5.2-744b.yml @@ -0,0 +1,74 @@ +# 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. + +# model config for GLM-5.2 - 744B (Mixture of Experts with Cross-Layer IndexShare) + +base_emb_dim: 6144 +base_num_query_heads: 64 +base_num_kv_heads: 64 +base_mlp_dim: 12288 +base_moe_mlp_dim: 2048 + +base_num_decoder_layers: 78 + +first_num_dense_layers: 3 +mlp_activations: ["silu","linear"] +vocab_size: 154880 +enable_dropout: false +logits_via_embedding: false +normalization_layer_epsilon: 1.0e-5 +num_experts: 256 +num_experts_per_tok: 8 +shared_experts: 1 +routed_scaling_factor: 2.5 +routed_score_func: "sigmoid" +routed_bias: true +float32_gate_logits: true +norm_topk_prob: true +decoder_block: "glm5" +dtype: "bfloat16" +weight_dtype: "bfloat16" + +# Multi-head Latent Attention (MLA) +attention_type: "mla" +attention: "dot_product" +q_lora_rank: 2048 +kv_lora_rank: 512 +qk_nope_head_dim: 192 +qk_rope_head_dim: 64 +v_head_dim: 256 + +# RoPE +mscale: 1.0 +rope_type: "yarn" +rope_max_timescale: 8000000 # "rope_theta": 8000000 +max_position_embeddings: 1048576 +original_max_position_embeddings: 4096 +rope_factor: 1.0 +beta_fast: 32 +beta_slow: 1 +rope_interleave: true +rope_truncate: true +rope_attention_scaling: false + +# Indexer for Dynamic Sparse Attention (DSA) +use_indexer: true +indexer_n_heads: 32 +indexer_head_dim: 128 +indexer_topk: 2048 + +# GLM-5.2 Cross-Layer IndexCache / IndexShare (Exact 78-layer HF topology: 3 Dense Full + 18x(3 Shared + 1 Full) + 3 Shared) +use_index_share: true +index_share_pattern: "FFFSSSFSSSFSSSFSSSFSSSFSSSFSSSFSSSFSSSFSSSFSSSFSSSFSSSFSSSFSSSFSSSFSSSFSSSFSSS" +prune_shared_indexers: true diff --git a/src/maxtext/configs/types.py b/src/maxtext/configs/types.py index db9a68d1cd..3192f5f134 100644 --- a/src/maxtext/configs/types.py +++ b/src/maxtext/configs/types.py @@ -240,6 +240,8 @@ class ProfilerType(str, Enum): "deepseek4-tiny", "deepseek4-284b", "deepseek-custom", + "glm5.1-744b", + "glm5.2-744b", "kimi-k2-1t", "gemma-7b", "gemma-2b", @@ -738,6 +740,17 @@ class AttentionIndexer(BaseModel): " during ties." ), ) + use_index_share: bool = Field( + False, description="Whether to enable cross-layer index cache sharing (GLM-5.2 IndexShare)." + ) + index_share_pattern: str = Field( + "FSSS", + description="Cross-layer pattern string for IndexShare (e.g. 'FSSS', 'F,S,S,S', 'FSFSS...').", + ) + prune_shared_indexers: bool = Field( + True, + description="Whether to prune indexer parameters on Shared (S) layers when use_index_share is enabled.", + ) class Llama4Attention(BaseModel): @@ -3563,7 +3576,7 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de self.tensors_to_offload = [t for t in tensors if getattr(self, t) == "offload"] if self.pipeline_parallel_layers == -1: - if self.decoder_block == DecoderBlockType.DEEPSEEK: + if self.decoder_block in (DecoderBlockType.DEEPSEEK, DecoderBlockType.GLM5): moe_layers = self.num_decoder_layers - self.first_num_dense_layers self.pipeline_parallel_layers = moe_layers else: @@ -3812,6 +3825,8 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de "when indexer loss is enabled (`indexer_loss_scaling_factor > 0.0`); otherwise the indexer " "short-circuits to select all tokens and no indexer loss is produced." ) + if not self.use_indexer and self.use_index_share: + raise ValueError("`use_index_share=True` requires `use_indexer=True`.") if not self.use_indexer and self.indexer_cutoff_threshold != RematLocation.REMAT: raise ValueError( f"Setting `indexer_cutoff_threshold='{self.indexer_cutoff_threshold}'` is only valid when " @@ -3914,9 +3929,9 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de if ( self.routed_bias and self.routed_bias_update_rate > 0.0 - and self.decoder_block not in (DecoderBlockType.DEEPSEEK, DecoderBlockType.DEEPSEEK4) + and self.decoder_block not in (DecoderBlockType.DEEPSEEK, DecoderBlockType.DEEPSEEK4, DecoderBlockType.GLM5) ): - raise ValueError("Loss-free load balancing is only supported for the DeepSeek decoder block.") + raise ValueError("Loss-free load balancing is only supported for the DeepSeek/GLM decoder block.") if not self.pure_nnx and self.routed_bias and self.decoder_block == DecoderBlockType.DEEPSEEK4: raise ValueError( "Auxiliary-loss-free routed bias for DeepSeek V4 is only supported in pure NNX mode. " diff --git a/src/maxtext/experimental/omni_poc/tests/custom_vision_projector_test.py b/src/maxtext/experimental/omni_poc/tests/custom_vision_projector_test.py index 3cadb389f4..58314cfc1e 100644 --- a/src/maxtext/experimental/omni_poc/tests/custom_vision_projector_test.py +++ b/src/maxtext/experimental/omni_poc/tests/custom_vision_projector_test.py @@ -12,9 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Tests for custom vision projector implementation in MaxText. - -""" +"""Tests for custom vision projector implementation in MaxText.""" import argparse import gc diff --git a/src/maxtext/inference/decode.py b/src/maxtext/inference/decode.py index 3034f5521e..418d6dd812 100644 --- a/src/maxtext/inference/decode.py +++ b/src/maxtext/inference/decode.py @@ -24,7 +24,7 @@ from maxtext.configs import pyconfig from maxtext.common import profiler -from maxtext.common.gcloud_stub import jetstream, is_decoupled +from maxtext.common.gcloud_stub import jetstream from maxtext.inference.maxengine import maxengine from maxtext.multimodal import processor as mm_processor from maxtext.multimodal import utils as mm_utils @@ -119,13 +119,6 @@ def main(argv: Sequence[str]) -> None: metadata = engine.get_tokenizer() tokenizer_model = engine.build_tokenizer(metadata) - token_params_is_stub = getattr(_token_params_ns, "_IS_STUB", False) - engine_api_is_stub = getattr(engine_api, "_IS_STUB", False) - if is_decoupled() and (token_params_is_stub or engine_api_is_stub): - raise RuntimeError( - "JetStream disabled by DECOUPLE_GCLOUD=TRUE or stubbed; decode requires the JetStream tokenizer. " - "Unset DECOUPLE_GCLOUD or install JetStream to run decode." - ) try: # TODO: update jetstream.engine.tokenizer_api.Tokenizer to maintain tokenizer state. diff --git a/src/maxtext/inference/maxengine/maxengine.py b/src/maxtext/inference/maxengine/maxengine.py index 62c5fe0345..ef859f9210 100644 --- a/src/maxtext/inference/maxengine/maxengine.py +++ b/src/maxtext/inference/maxengine/maxengine.py @@ -17,6 +17,7 @@ from collections import defaultdict from typing import Any, Callable import functools +import gc import os.path import uuid import warnings @@ -488,6 +489,7 @@ def _overlay(dst, src): nnx.replace_by_pure_dict(rest_state, rest_dict) self._nnx_rest_state = rest_state del nnx_model, loaded_rest_state, loaded_rest_dict, rest_dict + gc.collect() self.abstract_params = jax.tree.map( lambda x: jax.ShapeDtypeStruct(shape=x.shape, dtype=x.dtype, sharding=x.sharding) @@ -674,7 +676,7 @@ def _maybe_unstack_prefill_result_cache(self, cache): is_deepseek = ( getattr(self.model, "is_deepseek", False) or (hasattr(self.model, "decoder") and getattr(self.model.decoder, "is_deepseek", False)) - or (hasattr(self.config, "decoder_block") and str(self.config.decoder_block).lower() == "deepseek") + or (hasattr(self.config, "decoder_block") and str(self.config.decoder_block).lower() in ("deepseek", "glm5")) ) if is_deepseek: first_dense = self.config.first_num_dense_layers @@ -1496,11 +1498,12 @@ def copy(path, partial_cache, full_cache, annotations): ]: return full_cache # we don't even zero these out because we can mask them out. + ndim_diff = full_cache.ndim - len(annotations) batch_idx = -1 if "cache_batch" in annotations: - batch_idx = annotations.index("cache_batch") + batch_idx = annotations.index("cache_batch") + ndim_diff elif "cache_scale_batch" in annotations: - batch_idx = annotations.index("cache_scale_batch") + batch_idx = annotations.index("cache_scale_batch") + ndim_diff if batch_idx < 0: raise ValueError(f"Batch index {batch_idx=} shouldn't be less than zero for {path_key}, got {annotations=}") @@ -1527,7 +1530,7 @@ def copy(path, partial_cache, full_cache, annotations): ## copy prefill cachce full_cache = jax.lax.dynamic_update_index_in_dim(full_cache, partial_cache, slot, batch_idx) elif path_key == "cached_ar_lengths": - full_cache = full_cache.at[slot].set(0) + full_cache = full_cache.at[(slice(None),) * ndim_diff + (slot,)].set(0) elif path_key in [ "cached_prefill_key", "cached_prefill_value", @@ -1617,11 +1620,12 @@ def copy(path, partial_cache, full_cache, annotations): ]: return full_cache + ndim_diff = full_cache.ndim - len(annotations) batch_idx = -1 if "cache_batch" in annotations: - batch_idx = annotations.index("cache_batch") + batch_idx = annotations.index("cache_batch") + ndim_diff elif "cache_scale_batch" in annotations: - batch_idx = annotations.index("cache_scale_batch") + batch_idx = annotations.index("cache_scale_batch") + ndim_diff if batch_idx < 0: raise ValueError(f"Batch index {batch_idx=} shouldn't be less than zero for {path_key}, got {annotations=}") @@ -1645,7 +1649,7 @@ def copy(path, partial_cache, full_cache, annotations): full_cache = jax.lax.dynamic_update_index_in_dim(full_cache, partial_cache, slot, batch_idx) return full_cache elif path_key == "cached_ar_lengths": - return full_cache.at[slot].set(0) + return full_cache.at[(slice(None),) * ndim_diff + (slot,)].set(0) elif path_key in [ "cached_prefill_key", "cached_prefill_value", @@ -1755,11 +1759,12 @@ def copy(path, partial_cache, full_cache, annotations): ]: return full_cache # we don't even zero these out because we can mask them out. + ndim_diff = full_cache.ndim - len(annotations) batch_idx = -1 if "cache_batch" in annotations: - batch_idx = annotations.index("cache_batch") + batch_idx = annotations.index("cache_batch") + ndim_diff elif "cache_scale_batch" in annotations: - batch_idx = annotations.index("cache_scale_batch") + batch_idx = annotations.index("cache_scale_batch") + ndim_diff if batch_idx < 0: raise ValueError(f"Batch index {batch_idx=} shouldn't be less than zero for {path_key}, got {annotations=}") @@ -1770,10 +1775,14 @@ def copy(path, partial_cache, full_cache, annotations): if path_key == "cache_ar_segment_id": ### goal: zero this out in case there is existing data - zeros = jnp.zeros((1, self.config.max_target_length - self.config.max_prefill_predict_length), dtype=jnp.int32) + s = list(full_cache.shape) + s[batch_idx] = 1 + zeros = jnp.zeros(tuple(s), dtype=jnp.int32) return jax.lax.dynamic_update_index_in_dim(full_cache, zeros, slot, batch_idx) elif path_key == "cache_prefill_segment_id": - zeros = jnp.zeros((1, self.config.max_prefill_predict_length), dtype=jnp.int32) + s = list(full_cache.shape) + s[batch_idx] = 1 + zeros = jnp.zeros(tuple(s), dtype=jnp.int32) ## zero out in case prefill cache is too small to cover full_cache = jax.lax.dynamic_update_index_in_dim(full_cache, zeros, slot, batch_idx) # In case partial_cache is too small to slice at the given index, pad it with an extra seqlen @@ -1786,15 +1795,15 @@ def copy(path, partial_cache, full_cache, annotations): full_cache = jax.lax.dynamic_update_index_in_dim(full_cache, partial_cache, slot, batch_idx) return full_cache elif path_key == "cached_ar_lengths": - return full_cache.at[slot].set(0) + return full_cache.at[(slice(None),) * ndim_diff + (slot,)].set(0) elif path_key in [ "cached_prefill_key", "cached_prefill_value", "cached_prefill_key_scale", "cached_prefill_value_scale", ]: - seqlen_index = self.config.prefill_cache_axis_order.split(",").index("1") - start_indices = [0, 0, 0, 0] + seqlen_index = self.config.prefill_cache_axis_order.split(",").index("1") + ndim_diff + start_indices = [0] * full_cache.ndim start_indices[seqlen_index] = start_idx slice_size = list(partial_cache.shape) slice_size[seqlen_index] = seq_len @@ -1883,49 +1892,45 @@ def get_tokenizer(self) -> Any: When DECOUPLE_GCLOUD is FALSE we provide a clear error instead of failing cryptically on attribute access. """ - token_params_is_stub = getattr(_token_params_ns, "_IS_STUB", False) - engine_api_is_stub = getattr(engine_api, "_IS_STUB", False) - if is_decoupled() and (token_params_is_stub or engine_api_is_stub): - raise RuntimeError( - "JetStream disabled by DECOUPLE_GCLOUD=TRUE or stubbed; get_tokenizer is unsupported. " - "Unset DECOUPLE_GCLOUD or install JetStream to enable tokenizer functionality." - ) try: - # pyrefly: ignore[missing-attribute] - tokenizer_type_val = TokenizerType.DESCRIPTOR.values_by_name[ - self.config.tokenizer_type - ].number # pyrefly: ignore[missing-attribute] + tok_type = self.config.tokenizer_type + if hasattr(tok_type, "name"): + tok_name = tok_type.name.lower() + elif hasattr(tok_type, "value"): + tok_name = str(tok_type.value).lower() + else: + tok_name = str(tok_type).lower().rsplit(".", maxsplit=1)[-1] + tokenizer_val = getattr(TokenizerType, tok_name, None) + if tokenizer_val is None and hasattr(TokenizerType, "DESCRIPTOR"): + val_desc = getattr(TokenizerType.DESCRIPTOR, "values_by_name", {}) + if tok_name in val_desc: + tokenizer_val = val_desc[tok_name].number return TokenizerParameters( - path=self.config.tokenizer_path, # pyrefly: ignore[unexpected-keyword] - tokenizer_type=tokenizer_type_val, # pyrefly: ignore[unexpected-keyword] - access_token=self.config.hf_access_token, # pyrefly: ignore[unexpected-keyword] - use_chat_template=self.config.use_chat_template, # pyrefly: ignore[unexpected-keyword] - extra_ids=0, # pyrefly: ignore[unexpected-keyword] + path=self.config.tokenizer_path, + tokenizer_type=tokenizer_val if tokenizer_val is not None else tok_name, + access_token=self.config.hf_access_token, + use_chat_template=self.config.use_chat_template, + extra_ids=0, ) except KeyError as _: raise KeyError(f"Unsupported tokenizer type: {self.config.tokenizer_type}") from None def build_tokenizer(self, metadata: Any): # return type depends on JetStream """Return a tokenizer""" - token_params_is_stub = getattr(_token_params_ns, "_IS_STUB", False) - engine_api_is_stub = getattr(engine_api, "_IS_STUB", False) - if is_decoupled() and (token_params_is_stub or engine_api_is_stub): - raise RuntimeError( - "JetStream disabled by DECOUPLE_GCLOUD=TRUE or stubbed; build_tokenizer is unsupported. " - "Unset DECOUPLE_GCLOUD or install JetStream to enable tokenizer functionality." - ) - if metadata.tokenizer_type == TokenizerType.tiktoken: # pyrefly: ignore[missing-attribute] + tok_type = getattr(metadata, "tokenizer_type", None) + if tok_type in (getattr(TokenizerType, "tiktoken", None), "tiktoken"): return token_utils.TikToken(metadata) - elif metadata.tokenizer_type == TokenizerType.sentencepiece: # pyrefly: ignore[missing-attribute] + elif tok_type in (getattr(TokenizerType, "sentencepiece", None), "sentencepiece"): return token_utils.SentencePieceTokenizer(metadata) - elif metadata.tokenizer_type == TokenizerType.huggingface: # pyrefly: ignore[missing-attribute] + elif tok_type in (getattr(TokenizerType, "huggingface", None), "huggingface"): tokenizer_model = token_utils.HuggingFaceTokenizer(metadata) - if tokenizer_model.tokenizer.pad_token_id is None: - if tokenizer_model.tokenizer.unk_token_id is not None: - tokenizer_model.tokenizer.pad_token_id = tokenizer_model.tokenizer.unk_token_id - else: - print(f"Warning: setting pad_token_id to eos_token_id:{tokenizer_model.tokenizer.eos_token_id}") - tokenizer_model.tokenizer.pad_token_id = tokenizer_model.tokenizer.eos_token_id + tok = getattr(tokenizer_model, "tokenizer", tokenizer_model) + if hasattr(tok, "pad_token_id") and tok.pad_token_id is None: + if getattr(tok, "unk_token_id", None) is not None: + tok.pad_token_id = tok.unk_token_id + elif getattr(tok, "eos_token_id", None) is not None: + print(f"Warning: setting pad_token_id to eos_token_id:{tok.eos_token_id}") + tok.pad_token_id = tok.eos_token_id return tokenizer_model else: raise ValueError(f"Unsupported tokenizer type: {metadata.tokenizer_type}") diff --git a/src/maxtext/layers/attention_mla.py b/src/maxtext/layers/attention_mla.py index 3fc1a3c69d..4e92d64aaa 100644 --- a/src/maxtext/layers/attention_mla.py +++ b/src/maxtext/layers/attention_mla.py @@ -73,6 +73,7 @@ from maxtext.inference.kvcache import KVQuant from maxtext.utils.sharding import create_sharding from maxtext.utils.globals import EPS +from maxtext.utils import index_share_utils PLACEHOLDER_SEQ_LEN = 1 @@ -516,6 +517,8 @@ def mla_as_linen( mscale: float = 1.0, # scaling factor for softmax rope_factor: float = 40.0, # rotary embedding factor name: str | None = None, + is_shared_layer: bool = False, + served_group_size: int = 1, ): """A factory function to create an MLA as a Linen module. @@ -582,6 +585,8 @@ def mla_as_linen( mscale=mscale, rope_factor=rope_factor, name=name, + is_shared_layer=is_shared_layer, + served_group_size=served_group_size, metadata_fn=variable_to_logically_partitioned, abstract_init=False, ) @@ -654,6 +659,8 @@ def __init__( mscale: float = 1.0, # scaling factor for softmax rope_factor: float = 40.0, # rotary embedding factor name: str | None = None, + is_shared_layer: bool = False, + served_group_size: int = 1, rngs: Optional[nnx.Rngs] = None, ): """Initializes the MLA module. @@ -733,12 +740,26 @@ def __init__( # Initialize Indexer self.use_indexer = config.use_indexer - if self.use_indexer: + self.is_shared_layer = is_shared_layer + self.served_group_size = served_group_size + if getattr(config, "use_index_share", False): + pattern = index_share_utils.parse_index_share_pattern(config.index_share_pattern, config.num_decoder_layers) + self.is_full_tuple = tuple(role == "F" for role in pattern) + self.served_group_sizes_tuple = index_share_utils.get_served_group_sizes(pattern) + else: + self.is_full_tuple = None + self.served_group_sizes_tuple = None + is_pruned = ( + getattr(config, "use_index_share", False) + and getattr(config, "prune_shared_indexers", True) + and self.is_shared_layer + ) + if self.use_indexer and not is_pruned: # Need two versions of rope. # MLA applies yarn with interleave layout. # Indexer applies yarn with concatenate layout. indexer_rope = copy.copy(self.rotary_embedding) - indexer_rope.interleave = False + indexer_rope.interleave = getattr(config, "indexer_rope_interleave", config.rope_interleave) self.indexer = Indexer( config, rngs=rngs, @@ -1244,7 +1265,9 @@ def __call__( rope_kwargs: dict | None = None, kv_cache: Optional[Array] = None, attention_metadata: Optional[dict[str, Any]] = None, - ) -> tuple[Array, Optional[Array]]: + cached_indexer_state: Optional[Any] = None, + layer_idx: Optional[Any] = None, + ) -> tuple[Array, Optional[Array]] | tuple[Array, Optional[Array], Optional[Any]]: """Forward pass for MLA, reusing `AttentionOp` for the actual attention. Args: @@ -1259,10 +1282,10 @@ def __call__( bidirectional_mask: A mask for bidirectional attention, used in multimodal models. kv_cache: Optional key-value cache used when serving models with vLLM. attention_metadata: Optional attention-related metadata used when serving models with vLLM. + cached_indexer_state: Optional tuple (indexer_mask, topk_indices, indexer_score) from donor F-layer. Returns: - A tensor of shape [batch, length, embed_dim] containing the - MLA-attended outputs. + A tensor of shape [batch, length, embed_dim] containing the MLA-attended outputs. """ if model_mode == MODEL_MODE_PREFILL: inputs_q = self._maybe_shard_with_logical(inputs_q, self.prefill_input_axis_names) @@ -1288,33 +1311,83 @@ def __call__( # Indexer Logic indexer_mask = None + new_indexer_state = None if self.use_indexer: - # generate mask: with 0 and large negative, [b, 1, 1, q_len, kv_len] -> [b, q_len, kv_len] - attention_mask = self.attention_op.generate_attention_mask( - query, - key, - decoder_segment_ids, - model_mode, - previous_chunk, - bidirectional_mask, - segment_positions=inputs_positions, - ) - if attention_mask is not None: - attention_mask = attention_mask.squeeze(axis=(1, 2)) - # apply indexer, indexer_mask [b, q_len, kv_len] - indexer_mask, _, indexer_score = self.indexer( - inputs_q=inputs_q, - low_rank_q=low_rank_q, - inputs_kv=inputs_kv, - inputs_positions=inputs_positions, - attention_mask=attention_mask, - decoder_segment_ids=decoder_segment_ids, - previous_chunk=previous_chunk, - kv_cache=self.IndexerKVCache_0, - model_mode=model_mode, - ) + seq_len = key.shape[1] if key is not None else 0 + if seq_len <= self.config.indexer_topk: + indexer_mask, topk_indices, indexer_score = None, None, None + new_indexer_state = cached_indexer_state + else: + # generate mask: with 0 and large negative, [b, 1, 1, q_len, kv_len] -> [b, q_len, kv_len] + attention_mask = self.attention_op.generate_attention_mask( + query, + key, + decoder_segment_ids, + model_mode, + previous_chunk, + bidirectional_mask, + segment_positions=inputs_positions, + ) + if attention_mask is not None: + attention_mask = attention_mask.squeeze(axis=(1, 2)) + + if self.indexer is not None: + + def _run_full(_): + with jax.named_scope("glm_full_layer_indexer"): + mask, indices, score = self.indexer( + inputs_q=inputs_q, + low_rank_q=low_rank_q, + inputs_kv=inputs_kv, + inputs_positions=inputs_positions, + attention_mask=attention_mask, + decoder_segment_ids=decoder_segment_ids, + previous_chunk=previous_chunk, + kv_cache=self.IndexerKVCache_0, + model_mode=model_mode, + ) + mask = checkpoint_name(mask, "full_layer_indexer_mask") + indices = checkpoint_name(indices, "full_layer_topk_indices") + return mask, indices, score + + def _run_shared(_): + with jax.named_scope("glm_shared_layer_index_reuse"): + mask, indices, score = cached_indexer_state + mask = checkpoint_name(mask, "shared_layer_reused_mask") + indices = checkpoint_name(indices, "shared_layer_reused_indices") + return mask, indices, score + + if getattr(self.config, "use_index_share", False) and cached_indexer_state is not None: + if layer_idx is not None and self.is_full_tuple is not None: + is_full = jnp.array(self.is_full_tuple, dtype=jnp.bool_)[layer_idx] + indexer_mask, topk_indices, indexer_score = jax.lax.cond( + is_full, + _run_full, + _run_shared, + operand=None, + ) + elif self.is_shared_layer: + indexer_mask, topk_indices, indexer_score = _run_shared(None) + else: + indexer_mask, topk_indices, indexer_score = _run_full(None) + else: + indexer_mask, topk_indices, indexer_score = _run_full(None) + + new_indexer_state = (indexer_mask, topk_indices, indexer_score) + else: + indexer_mask, topk_indices, indexer_score = None, None, None + new_indexer_state = cached_indexer_state + + if indexer_mask is not None and self.config.indexer_loss_scaling_factor > 0.0 and indexer_score is not None: + loss_scale = self.config.indexer_loss_scaling_factor + if getattr(self.config, "use_index_share", False): + group_size = ( + jnp.array(self.served_group_sizes_tuple, dtype=jnp.float32)[layer_idx] + if layer_idx is not None and self.served_group_sizes_tuple is not None + else float(self.served_group_size) + ) + loss_scale = loss_scale / jnp.maximum(group_size, 1.0) - if indexer_mask is not None and self.config.indexer_loss_scaling_factor > 0.0: indexer_loss = self.calculate_indexer_loss( indexer_score=indexer_score, query=query, @@ -1322,7 +1395,7 @@ def __call__( attention_mask=attention_mask, indexer_mask=indexer_mask, sparse_loss=self.config.indexer_sparse_training, - scaling_factor=self.config.indexer_loss_scaling_factor, + scaling_factor=loss_scale, ) self.indexer_loss = indexer_losses(indexer_loss) @@ -1347,4 +1420,6 @@ def __call__( out_sharding = create_sharding(self.mesh, out_logical_name) out = self.out_projection(out, out_sharding=out_sharding) out = checkpoint_name(out, "out_proj") + if getattr(self.config, "use_index_share", False): + return out, kv_cache, new_indexer_state return out, kv_cache diff --git a/src/maxtext/layers/attentions.py b/src/maxtext/layers/attentions.py index a2f4b48afd..361c4f547a 100644 --- a/src/maxtext/layers/attentions.py +++ b/src/maxtext/layers/attentions.py @@ -1067,6 +1067,9 @@ def init_kv_caches(self, inputs_kv_shape: Tuple): # KVCache. placeholder_seq_len = 1 + key_head_size = getattr(self, "qk_head_dim", self.head_dim) + value_head_size = getattr(self, "v_head_dim", self.head_dim) + return kvcache.KVCache( max_prefill_length=self.max_prefill_predict_length, max_target_length=self.max_target_length, @@ -1075,8 +1078,8 @@ def init_kv_caches(self, inputs_kv_shape: Tuple): value_seq_len=placeholder_seq_len, key_heads=self.num_kv_heads, value_heads=self.num_kv_heads, - key_head_size=self.head_dim, - value_head_size=self.head_dim, + key_head_size=key_head_size, + value_head_size=value_head_size, dtype=self.dtype, kv_quant=self.kv_quant, prefill_cache_axis_order=self.prefill_cache_axis_order, diff --git a/src/maxtext/layers/decoders.py b/src/maxtext/layers/decoders.py index 648392738e..36df15e955 100644 --- a/src/maxtext/layers/decoders.py +++ b/src/maxtext/layers/decoders.py @@ -50,6 +50,7 @@ gemma3, gemma4, gemma4_small, + glm5, gpt3, gpt_oss, llama2, @@ -454,6 +455,11 @@ def get_decoder_layers(self): deepseek.DeepSeekDenseLayerToLinen, deepseek.DeepSeekMoELayerToLinen, ] + case DecoderBlockType.GLM5: + return [ + glm5.GLMDenseLayerToLinen, + glm5.GLMMoELayerToLinen, + ] case DecoderBlockType.DEEPSEEK4: return ( [deepseek4.DeepSeek4ScannableBlockToLinen] if self.config.scan_layers else [deepseek4.DeepSeek4LayerToLinen] @@ -529,6 +535,7 @@ def get_scannable(normal_cls, scannable_cls): DecoderBlockType.SIMPLE: [simple_layer.SimpleDecoderLayer], DecoderBlockType.SIMPLE_MLP: [simple_layer.SimpleMlpDecoderLayer], DecoderBlockType.DEEPSEEK: [deepseek.DeepSeekDenseLayer, deepseek.DeepSeekMoELayer], + DecoderBlockType.GLM5: [glm5.GLMDenseLayer, glm5.GLMMoELayer], DecoderBlockType.LLAMA4: get_scannable(llama4.Llama4DecoderLayer, llama4.Llama4ScannableBlock), DecoderBlockType.OLMO3: get_scannable(olmo3.Olmo3DecoderLayer, olmo3.Olmo3ScannableBlock), DecoderBlockType.ENVY: get_scannable(envy.EnvyDecoderLayer, envy.EnvyScannableBlock), @@ -569,7 +576,11 @@ def map_fn(path, value): def _build_nnx_pipeline_stage(self, decoder_blocks, rngs): """Creates a single NNX pipeline stage module.""" cfg = self.config - base_stage_cls = decoder_blocks[1] if cfg.decoder_block == DecoderBlockType.DEEPSEEK else decoder_blocks[0] + base_stage_cls = ( + decoder_blocks[1] + if cfg.decoder_block in (DecoderBlockType.DEEPSEEK, DecoderBlockType.GLM5) + else decoder_blocks[0] + ) if cfg.num_layers_per_pipeline_stage == 1: return base_stage_cls(config=cfg, mesh=self.mesh, quant=self.quant, model_mode=self.model_mode, rngs=rngs) @@ -585,7 +596,7 @@ def get_pipeline_stage_module(self, decoder_blocks): """get pipeline stage module""" def get_layer_to_pipeline(blocks, cfg): - if cfg.decoder_block == DecoderBlockType.DEEPSEEK: + if cfg.decoder_block in (DecoderBlockType.DEEPSEEK, DecoderBlockType.GLM5): return blocks[1] # return the sparse block else: return blocks[0] @@ -627,6 +638,7 @@ def get_norm_layer(self, num_features: int): DecoderBlockType.MIXTRAL, DecoderBlockType.DEEPSEEK, DecoderBlockType.DEEPSEEK4, + DecoderBlockType.GLM5, DecoderBlockType.GEMMA, DecoderBlockType.GEMMA2, DecoderBlockType.GEMMA3, @@ -897,8 +909,8 @@ def __call__( if cfg.pipeline_fsdp_ag_once or cfg.pipeline_fsdp_ag_per_repeat else None ) - if cfg.decoder_block == DecoderBlockType.DEEPSEEK: - assert len(RemattedBlockLayers) == 2, "Scanned layers must have a length of 2 using deepseek." + if cfg.decoder_block in (DecoderBlockType.DEEPSEEK, DecoderBlockType.GLM5): + assert len(RemattedBlockLayers) == 2, "Scanned layers must have a length of 2 using deepseek/glm." dense_layer = RemattedBlockLayers[0] moe_layer = RemattedBlockLayers[1] num_moe_layers = cfg.num_decoder_layers - cfg.first_num_dense_layers @@ -943,8 +955,8 @@ def __call__( )(y, *broadcast_args) else: if cfg.scan_layers: - if cfg.decoder_block == DecoderBlockType.DEEPSEEK: - assert len(RemattedBlockLayers) == 2, "Scanned layers must have a length of 2 using deepseek." + if cfg.decoder_block in (DecoderBlockType.DEEPSEEK, DecoderBlockType.GLM5): + assert len(RemattedBlockLayers) == 2, "Scanned layers must have a length of 2 using deepseek/glm." layer_call_kwargs = { "previous_chunk": previous_chunk, "slot": slot, @@ -1149,8 +1161,8 @@ def __call__( **layer_kwargs, )(y, *current_broadcast_args) else: - if cfg.decoder_block == DecoderBlockType.DEEPSEEK: - assert len(RemattedBlockLayers) == 2, "Unscanned layers must have a length of 2 using deepseek." + if cfg.decoder_block in (DecoderBlockType.DEEPSEEK, DecoderBlockType.GLM5): + assert len(RemattedBlockLayers) == 2, "Unscanned layers must have a length of 2 using deepseek/glm." dense_layer = RemattedBlockLayers[0] moe_layer = RemattedBlockLayers[1] diff --git a/src/maxtext/layers/embeddings.py b/src/maxtext/layers/embeddings.py index a744ea041a..fc7d6eda7c 100644 --- a/src/maxtext/layers/embeddings.py +++ b/src/maxtext/layers/embeddings.py @@ -2044,7 +2044,7 @@ def __call__( self, inputs: jnp.ndarray, position: jnp.ndarray, - unsqueeze_dim: int | None = 1, + unsqueeze_dim: int | None = 2, reverse: bool = False, ) -> jnp.ndarray: """Applies interleaved Rotary Position Embedding to the inputs. @@ -2113,6 +2113,8 @@ def _apply_rotary_pos_emb( # Insert an expansion dimension to align the frequency tensors (e.g., [B, S, D]) # with the attention head axes of the input tensor (e.g., [B, S, H, D]). if unsqueeze_dim is not None: + if unsqueeze_dim < 0: + unsqueeze_dim = cos.ndim + unsqueeze_dim + 1 cos = jnp.expand_dims(cos, axis=unsqueeze_dim) sin = jnp.expand_dims(sin, axis=unsqueeze_dim) diff --git a/src/maxtext/layers/linears.py b/src/maxtext/layers/linears.py index 8e14d6d862..b442dc445a 100644 --- a/src/maxtext/layers/linears.py +++ b/src/maxtext/layers/linears.py @@ -563,6 +563,7 @@ def get_norm_layer(self, num_features: int): DecoderBlockType.GEMMA3, DecoderBlockType.QWEN3, DecoderBlockType.DEEPSEEK, + DecoderBlockType.GLM5, DecoderBlockType.LLAMA4, ): return functools.partial(normalizations.RMSNorm, num_features=num_features) diff --git a/src/maxtext/layers/moe.py b/src/maxtext/layers/moe.py index da9e86e320..7c67fb5c57 100644 --- a/src/maxtext/layers/moe.py +++ b/src/maxtext/layers/moe.py @@ -383,7 +383,7 @@ def __call__(self, inputs: jax.Array, _initializing: bool = False) -> Tuple[jax. output = linears._convert_to_activation_function(self.score_func)(output) # NOTE: deepseek2 has a different pattern - if self.model_name.startswith(("deepseek3", "deepseek4")): + if self.model_name.startswith(("deepseek3", "deepseek4", "glm5")): pre_bias_logits = output if self.use_bias: @@ -737,7 +737,7 @@ def get_topk(self, gate_logits, pre_bias_logits, rngs=None, input_ids=None): top_k_indices = tid2eid_int[input_ids.astype(jnp.int32)] top_k_weights = jnp.take_along_axis(pre_bias_logits, top_k_indices, axis=-1) # NOTE: deepseek2 has a different pattern - elif self.config.model_name.startswith(("deepseek3", "deepseek4")): + elif self.config.model_name.startswith(("deepseek3", "deepseek4", "glm5")): top_k_weights, top_k_indices = self.deepseek_routing(gate_logits, pre_bias_logits) elif self.config.decoder_block == ctypes.DecoderBlockType.GEMMA4: router_probs = jax.nn.softmax(gate_logits.astype(jnp.float32), axis=-1) @@ -746,7 +746,11 @@ def get_topk(self, gate_logits, pre_bias_logits, rngs=None, input_ids=None): else: top_k_weights, top_k_indices = jax.lax.top_k(gate_logits, self.num_experts_per_tok) - if self.config.decoder_block in (ctypes.DecoderBlockType.DEEPSEEK, ctypes.DecoderBlockType.DEEPSEEK4): + if self.config.decoder_block in ( + ctypes.DecoderBlockType.DEEPSEEK, + ctypes.DecoderBlockType.DEEPSEEK4, + ctypes.DecoderBlockType.GLM5, + ): top_k_weights = self.deepseek_scale_weights(top_k_weights) else: if self.config.decoder_block not in (ctypes.DecoderBlockType.LLAMA4, ctypes.DecoderBlockType.GEMMA4): @@ -852,7 +856,8 @@ def apply_ffn_activation(self, layer_w0, layer_w1): glu = jnp.multiply(layer_w0, layer_act) intermediate_layer = jnp.multiply(glu, (layer_w1 + 1)) elif ( - self.config.decoder_block in (ctypes.DecoderBlockType.DEEPSEEK, ctypes.DecoderBlockType.DEEPSEEK4) + self.config.decoder_block + in (ctypes.DecoderBlockType.DEEPSEEK, ctypes.DecoderBlockType.DEEPSEEK4, ctypes.DecoderBlockType.GLM5) and self.config.mlp_activations_limit > 0.0 ): # DeepSeek V4 uses bounds to clip the SwiGLU activations @@ -1625,7 +1630,7 @@ def get_routed_moe_shardings(is_batch_sharded_by_expert, has_input_ids): gate_logits_pspec = self._logical_to_mesh_axes((batch_logical_axis, "activation_norm_length", None)) # NOTE: deepseek2 has a different pattern - if self.config.model_name.startswith(("deepseek3", "deepseek4")): + if self.config.model_name.startswith(("deepseek3", "deepseek4", "glm5")): pre_bias_logits_pspec = self._logical_to_mesh_axes((batch_logical_axis, "activation_norm_length", None)) else: # pre_bias_logits is None for non-deepseek3/4 models, including deepseek2 @@ -2425,7 +2430,7 @@ def sparse_matmul_route_and_compute( gate_logits_logical_axes = (batch_logical_axis, "activation_norm_length", None) pre_bias_logits_logical_axes = ( (batch_logical_axis, "activation_norm_length", None) - if self.config.model_name.startswith(("deepseek3", "deepseek4")) + if self.config.model_name.startswith(("deepseek3", "deepseek4", "glm5")) else None ) inputs = self._maybe_shard_with_pspec(inputs, input_partition_pspec, logical_axes=input_logical_axes) @@ -2746,7 +2751,7 @@ def dense_matmul( # gate_logits: batch, length, expert gate_logits = self._maybe_shard_with_logical(gate_logits, ("activation_batch_moe", "activation_length_moe", None)) # NOTE: deepseek2 has a different pattern - if self.config.model_name.startswith(("deepseek3", "deepseek4")): + if self.config.model_name.startswith(("deepseek3", "deepseek4", "glm5")): # pre_bias_logits is None for non-deepseek3/4 models, including deepseek2 pre_bias_logits = self._maybe_shard_with_logical( pre_bias_logits, ("activation_batch_moe", "activation_length_moe", None) @@ -2824,10 +2829,7 @@ def dense_matmul( mlp_down_einsum = "EBCH,EHM -> EBCM" output_einsum = "EBCM,BSEC -> BSM" else: - # TODO(b/425930507): Try replacing `softmax_probs` with padded weights - # and verify with decode acc tests. - softmax_probs = jax.nn.softmax(gate_logits.astype(jnp.float32), axis=-1).astype(self.dtype) - dispatch_mask, combine_mask = self.generate_masks_subgroup(top_k_indices, softmax_probs) + dispatch_mask, combine_mask = self.generate_masks_subgroup(top_k_indices, weights) if self.get_context_autoregressive_parallelism_size() > 0 and cp == 1: mask_axes = ( "activation_norm_length_moe", diff --git a/src/maxtext/layers/nnx_decoders.py b/src/maxtext/layers/nnx_decoders.py index 1a9fdd48b0..b1ee2cee7f 100644 --- a/src/maxtext/layers/nnx_decoders.py +++ b/src/maxtext/layers/nnx_decoders.py @@ -15,6 +15,7 @@ """Module for decoder layers""" # pylint: disable=arguments-differ # pylint: disable=no-name-in-module +# pylint: disable=unbalanced-tuple-unpacking import functools import inspect @@ -54,6 +55,7 @@ gemma3, gemma4, gemma4_small, + glm5, gpt3, gpt_oss, llama2, @@ -432,7 +434,7 @@ def __init__( ) self.scanned_layers = None - self.is_deepseek = self.config.decoder_block == DecoderBlockType.DEEPSEEK + self.is_deepseek = self.config.decoder_block in (DecoderBlockType.DEEPSEEK, DecoderBlockType.GLM5) self.is_deepseek4 = self.config.decoder_block == DecoderBlockType.DEEPSEEK4 self.is_gemma3 = self.config.decoder_block == DecoderBlockType.GEMMA3 self.is_gemma4 = self.config.decoder_block == DecoderBlockType.GEMMA4 @@ -755,9 +757,10 @@ def _init_sequential_deepseek(self, decoder_block_classes, rngs): config = self.config dense_cls, moe_cls = decoder_block_classes for i in range(config.first_num_dense_layers): - self._create_and_register_layer(dense_cls, rngs, "dense_layers", i) + self._create_and_register_layer(dense_cls, rngs, "dense_layers", i, layer_idx=i) for i in range(config.num_decoder_layers - config.first_num_dense_layers): - self._create_and_register_layer(moe_cls, rngs, "moe_layers", i) + global_idx = config.first_num_dense_layers + i + self._create_and_register_layer(moe_cls, rngs, "moe_layers", i, layer_idx=global_idx) def _init_sequential_generic(self, decoder_block_classes, rngs): """Initializes sequential generic decoder layers with per-architecture layer_kwargs.""" @@ -959,6 +962,13 @@ def _apply_layers_sequentially( (final_carry, updated_layers, returned_kv_stacked) otherwise. """ if length == 0: + if getattr(self.config, "use_index_share", False): + return ( + x_in, + layers, + kv_caches_stacked if kv_caches_stacked is not None else None, + kwargs.get("cached_indexer_state", None), + ) return ( x_in, layers, @@ -987,10 +997,26 @@ def _extract_matching_state(template, full): updated_graphdef = [graphdef] use_kv = kv_caches_stacked is not None + is_index_share = getattr(self.config, "use_index_share", False) + cached_indexer_state = kwargs.get("cached_indexer_state", None) + start_layer_idx = kwargs.get("start_layer_idx", 0) + + if is_index_share: + cached_indexer_state = kwargs.get("cached_indexer_state", None) + if cached_indexer_state is None: + batch, seq_len = x_in.shape[0], x_in.shape[1] + topk = getattr(self.config, "indexer_topk", 2048) + dummy_mask = jnp.zeros((batch, seq_len, seq_len), dtype=jnp.bfloat16) + dummy_indices = jnp.zeros((batch, seq_len, topk), dtype=jnp.int32) + dummy_score = jnp.zeros((batch, seq_len, seq_len), dtype=jnp.float32) + cached_indexer_state = (dummy_mask, dummy_indices, dummy_score) + init_scan_carry = (x_in, cached_indexer_state, start_layer_idx) + else: + init_scan_carry = x_in def layer_fn(carry, scanned_vars): # Ensure metadata rank matches the sliced values - scanned_vars = maxtext_utils_nnx.nnx_remove_scan_axis(scanned_vars, "layers") + scanned_vars = maxtext_utils_nnx.nnx_remove_scan_axis(scanned_vars, metadata_axis_name) # Unpack the sliced variables for THIS layer if use_kv: @@ -1012,14 +1038,28 @@ def layer_fn(carry, scanned_vars): if kv_cache_layer is not None: call_kwargs["kv_cache"] = kv_cache_layer - layer_out = layer(carry, *args, **call_kwargs) + if is_index_share: + y_in, current_cached_indexer, lyr_idx = carry + call_kwargs["cached_indexer_state"] = current_cached_indexer + call_kwargs["layer_idx"] = lyr_idx + else: + y_in = carry + + layer_out = layer(y_in, *args, **call_kwargs) if isinstance(layer_out, tuple): - new_carry = layer_out[0] + new_carry_y = layer_out[0] updated_kv = layer_out[1] if len(layer_out) > 1 else None + new_indexer_state = layer_out[2] if len(layer_out) > 2 else None else: - new_carry = layer_out + new_carry_y = layer_out updated_kv = None + new_indexer_state = None + + if is_index_share: + new_carry = (new_carry_y, new_indexer_state, lyr_idx + 1) + else: + new_carry = new_carry_y # Extract the updated state to return it if dynamic_graph_init: @@ -1052,7 +1092,7 @@ def layer_fn(carry, scanned_vars): # kv_caches_stacked is actually the original kv_caches list in this new flow kv_caches_list = kv_caches_stacked - current_carry = x_in + current_carry = init_scan_carry for i in range(length): # Statically slice the parameters and state for this layer @@ -1067,16 +1107,30 @@ def layer_fn(carry, scanned_vars): # Update the list in-place (mutates the list passed by reference) kv_caches_list[i] = updated_kv + if is_index_share: + final_carry, out_indexer_state, _ = current_carry + else: + final_carry = current_carry + out_indexer_state = None + # We don't need to rebuild scanned_state or return it because during # inference with vLLM, parameters do not change and we don't need intermediates. - return current_carry, layers, None + if is_index_share: + return final_carry, layers, None, out_indexer_state + return final_carry, layers, None else: params = maxtext_utils_nnx.nnx_ensure_scan_leading_axis(params, length) state = maxtext_utils_nnx.nnx_ensure_scan_leading_axis(state, length) - final_carry, scanned_state = jax.lax.scan(layer_fn_wrapped, x_in, (params, state), unroll=unroll) + scan_res_carry, scanned_state = jax.lax.scan(layer_fn_wrapped, init_scan_carry, (params, state), unroll=unroll) returned_kv_stacked = None + if is_index_share: + final_carry, out_indexer_state, _ = scan_res_carry + else: + final_carry = scan_res_carry + out_indexer_state = None + # Move the scan axis to each variable's param_scan_axis and restore its name # in the sharding metadata. jax.lax.scan emits it at position 0. scanned_state = maxtext_utils_nnx.nnx_add_and_sync_scan_axis(scanned_state, metadata_axis_name) @@ -1093,6 +1147,8 @@ def layer_fn(carry, scanned_vars): nnx.update(layers, clean_state) out_layers = layers + if is_index_share: + return final_carry, out_layers, returned_kv_stacked if use_kv else None, out_indexer_state return final_carry, out_layers, returned_kv_stacked if use_kv else None def get_decoder_layers(self): @@ -1123,6 +1179,7 @@ def get_deepseek(): DecoderBlockType.SIMPLE: [simple_layer.SimpleDecoderLayer], DecoderBlockType.SIMPLE_MLP: [simple_layer.SimpleMlpDecoderLayer], DecoderBlockType.DEEPSEEK: get_deepseek(), + DecoderBlockType.GLM5: [glm5.GLMDenseLayer, glm5.GLMMoELayer], DecoderBlockType.DEEPSEEK4: get_scannable(deepseek4.DeepSeek4DecoderLayer, deepseek4.DeepSeek4ScannableBlock), DecoderBlockType.GPT_OSS: get_scannable(gpt_oss.GptOssDecoderLayer, gpt_oss.GptOssScannableBlock), DecoderBlockType.QWEN3_NEXT: get_scannable(qwen3.Qwen3NextDecoderLayer, qwen3.Qwen3NextScannableBlock), @@ -1276,6 +1333,7 @@ def get_norm_layer(self, num_features: int, rngs: nnx.Rngs): DecoderBlockType.MIXTRAL, DecoderBlockType.DEEPSEEK, DecoderBlockType.DEEPSEEK4, + DecoderBlockType.GLM5, DecoderBlockType.GEMMA, DecoderBlockType.GEMMA2, DecoderBlockType.GEMMA3, @@ -1787,6 +1845,28 @@ def __call__( *layer_args, **common_kwargs, ) + elif getattr(cfg, "use_index_share", False): + y, self.dense_layers, _, cached_indexer_state = self._apply_layers_sequentially( + self.dense_layers, + y, + *layer_args, + length=cfg.first_num_dense_layers, + start_layer_idx=0, + cached_indexer_state=None, + **layer_kwargs, + ) + + num_moe = cfg.num_decoder_layers - cfg.first_num_dense_layers + + y, self.moe_layers, _, _ = self._apply_layers_sequentially( + self.moe_layers, + y, + *layer_args, + length=num_moe, + start_layer_idx=cfg.first_num_dense_layers, + cached_indexer_state=cached_indexer_state, + **layer_kwargs, + ) else: y, self.dense_layers, _ = self._apply_layers_sequentially( self.dense_layers, @@ -1892,17 +1972,27 @@ def pure_layer_fn(graphdef_in, state_in, y_in, kv_in): state_in, ) merged_layer = nnx.merge(graphdef_in, state_in) - out_y, out_kv = merged_layer(y_in, *layer_args, kv_cache=kv_in, **layer_kwargs) + out = merged_layer(y_in, *layer_args, kv_cache=kv_in, **layer_kwargs) + if getattr(cfg, "use_index_share", False): + out_y, out_kv, out_indexer_cache = out + else: + out_y, out_kv = out + out_indexer_cache = None state_out = nnx.state(merged_layer) if dynamic_graph_init: new_graphdef, _, _ = nnx.split(merged_layer, nnx.Param, ...) + if getattr(cfg, "use_index_share", False): + return out_y, out_kv, out_indexer_cache, state_out, new_graphdef return out_y, out_kv, state_out, new_graphdef else: + if getattr(cfg, "use_index_share", False): + return out_y, out_kv, out_indexer_cache, state_out, graphdef_in return out_y, out_kv, state_out, graphdef_in checkpointed_fn = jax.checkpoint(pure_layer_fn, policy=policy, prevent_cse=prevent_cse) + cached_indexer_state = None for lyr in range(cfg.num_decoder_layers): if self.is_deepseek: if lyr < cfg.first_num_dense_layers: @@ -1948,11 +2038,18 @@ def pure_layer_fn(graphdef_in, state_in, y_in, kv_in): ) if input_tokens is not None: layer_kwargs["decoder_input_tokens"] = input_tokens + if getattr(cfg, "use_index_share", False): + layer_kwargs["cached_indexer_state"] = cached_indexer_state if cfg.remat_policy != "none": - y, kv_cache, new_state, new_graphdef = checkpointed_fn(graphdef, state, y, kv_cache) + res = checkpointed_fn(graphdef, state, y, kv_cache) + else: + res = pure_layer_fn(graphdef, state, y, kv_cache) + + if getattr(cfg, "use_index_share", False): + y, kv_cache, cached_indexer_state, new_state, new_graphdef = res else: - y, kv_cache, new_state, new_graphdef = pure_layer_fn(graphdef, state, y, kv_cache) + y, kv_cache, new_state, new_graphdef = res if dynamic_graph_init: new_layer = nnx.merge(new_graphdef, new_state) diff --git a/src/maxtext/layers/nnx_wrappers.py b/src/maxtext/layers/nnx_wrappers.py index e204502cb2..3bfe98a675 100644 --- a/src/maxtext/layers/nnx_wrappers.py +++ b/src/maxtext/layers/nnx_wrappers.py @@ -29,7 +29,17 @@ from flax.nnx import variablelib from flax.nnx.bridge import module as bdg_module from flax.nnx.module import Module -from flax.nnx import Pytree + +try: + from flax.nnx import Pytree +except ImportError: + try: + from flax.nnx import PyTree as Pytree + except ImportError: + try: + from flax.nnx.object import Object as Pytree + except ImportError: + Pytree = object from flax.nnx.rnglib import Rngs import jax from jax import tree_util as jtu diff --git a/src/maxtext/models/deepseek.py b/src/maxtext/models/deepseek.py index 0ad8978e7f..847b434507 100644 --- a/src/maxtext/models/deepseek.py +++ b/src/maxtext/models/deepseek.py @@ -140,7 +140,8 @@ def __init__( self.engram = None # DeepSeek V4 natively overrides this block with CompressedAttention. - if self.config.decoder_block != DecoderBlockType.DEEPSEEK4: + # GLM-5.2 overrides this in GLMGenericLayer with IndexShare configuration. + if self.config.decoder_block not in (DecoderBlockType.DEEPSEEK4, DecoderBlockType.GLM5): self.self_attention = attention_mla.MLA( config=self.config, num_query_heads=self.config.num_query_heads, diff --git a/src/maxtext/models/glm5.py b/src/maxtext/models/glm5.py new file mode 100644 index 0000000000..3d923c9d11 --- /dev/null +++ b/src/maxtext/models/glm5.py @@ -0,0 +1,392 @@ +# 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. + +"""GLM model definitions (GLM-5.1 & GLM-5.2 with Cross-Layer IndexShare).""" +# pylint: disable=arguments-differ +# pylint: disable=no-name-in-module + +from typing import Optional + +import absl.logging +from flax import nnx +import jax +from jax.ad_checkpoint import checkpoint_name +import jax.numpy as jnp +from jax.sharding import Mesh +from maxtext.common.common_types import AttentionType, Config, HyperConnectionType +from maxtext.layers import attention_mla +from maxtext.layers import initializers +from maxtext.layers import linears +from maxtext.layers import moe +from maxtext.layers import nnx_wrappers +from maxtext.layers import quantizations +from maxtext.models import deepseek +from maxtext.utils import index_share_utils + + +class GLMGenericLayer(deepseek.DeepSeekGenericLayer): # pylint: disable=abstract-method + """Generic GLM layer with Multi-Head Latent Attention and IndexShare support.""" + + def __init__( + self, + config: Config, + model_mode: str, + mesh: Mesh, + rngs: nnx.Rngs, + quant: Optional[quantizations.AqtQuantization] = None, + layer_idx: int = -1, + ) -> None: + super().__init__(config, model_mode, mesh, rngs, quant, layer_idx) + + # GLM-5.2 Cross-Layer IndexShare Role Resolution + self.is_index_share_enabled = getattr(config, "use_index_share", False) + self.is_shared_layer = False + self.served_group_size = 1 + if self.is_index_share_enabled and layer_idx >= 0: + pattern = index_share_utils.parse_index_share_pattern(config.index_share_pattern, config.num_decoder_layers) + self.is_shared_layer = index_share_utils.is_shared_layer(layer_idx, pattern) + self.served_group_size = index_share_utils.get_served_group_sizes(pattern)[layer_idx] + if layer_idx == 0 and jax.process_index() == 0: + num_f = pattern.count("F") + num_s = pattern.count("S") + absl.logging.info( + "[GLM-5.2 IndexShare Active] Total layers: %d | Pattern: %s | Full (F) layers with active indexers: %d | " + "Shared (S) layers with pruned indexers: %d (Pruned %.1f%% indexer compute/parameters)", + config.num_decoder_layers, + config.index_share_pattern, + num_f, + num_s, + num_s / config.num_decoder_layers * 100, + ) + + # Re-initialize MLA with GLM-specific IndexShare configuration + self.self_attention = attention_mla.MLA( + config=self.config, + num_query_heads=self.config.num_query_heads, + num_kv_heads=self.config.num_kv_heads, + head_dim=self.config.head_dim, + max_target_length=self.config.max_target_length, + max_prefill_predict_length=self.config.max_prefill_predict_length, + attention_kernel=self.config.attention, + attention_type=AttentionType(self.config.attention_type), + inputs_q_shape=self.dummy_inputs_shape, + inputs_kv_shape=self.dummy_inputs_shape, + mesh=mesh, + dtype=self.config.dtype, + weight_dtype=self.config.weight_dtype, + dropout_rate=self.config.dropout_rate, + name="self_attention", + quant=quant, + kv_quant=quantizations.configure_kv_quant(self.config), + q_lora_rank=self.config.q_lora_rank, + kv_lora_rank=self.config.kv_lora_rank, + qk_nope_head_dim=self.config.qk_nope_head_dim, + qk_rope_head_dim=self.config.qk_rope_head_dim, + v_head_dim=self.config.v_head_dim, + max_position_embeddings=self.config.max_position_embeddings, + original_max_position_embeddings=self.config.original_max_position_embeddings, + mscale=self.config.mscale, + rope_factor=self.config.rope_factor, + model_mode=model_mode, + rngs=rngs, + attn_logits_soft_cap=self.config.attn_logits_soft_cap, + is_shared_layer=self.is_shared_layer, + served_group_size=self.served_group_size, + ) + + def attention_op( + self, + x, + decoder_segment_ids, + decoder_positions, + deterministic, + model_mode, + previous_chunk=None, + slot: None | int = None, + cached_indexer_state=None, + layer_idx=None, + ): + """Executes the attention layer and passes cached indexer state.""" + attn_out = self.self_attention( + x, + x, + decoder_positions, + decoder_segment_ids=decoder_segment_ids, + deterministic=deterministic, + model_mode=model_mode, + out_sharding=self.out_sharding, + previous_chunk=previous_chunk, + slot=slot, + cached_indexer_state=cached_indexer_state, + layer_idx=layer_idx, + ) + attention_result = attn_out[0] + if self.is_index_share_enabled: + new_indexer_state = attn_out[2] if len(attn_out) > 2 else None + return self.with_logical_constraint(attention_result), new_indexer_state + else: + return self.with_logical_constraint(attention_result), None + + def post_process(self, layer_output, load_balance_loss, moe_bias_updates, kv_cache=None, cached_indexer_state=None): + """Post-processing with IndexShare state pass-through.""" + if self.config.load_balance_loss_weight > 0.0 and load_balance_loss is not None: + self.sow(nnx.Intermediate, "moe_lb_loss", load_balance_loss) + + if self.config.routed_bias and self.config.routed_bias_update_rate > 0.0 and moe_bias_updates is not None: + self.sow(nnx.Intermediate, "moe_bias_updates", moe_bias_updates) + + if getattr(self.config, "record_internal_nn_metrics", False): + self.sow(nnx.Intermediate, "activation_mean", jnp.mean(layer_output)) + self.sow(nnx.Intermediate, "activation_stdev", jnp.std(layer_output)) + self.sow( + nnx.Intermediate, + "activation_fraction_zero", + jnp.sum(layer_output == 0) / jnp.size(layer_output), + ) + + if self.is_index_share_enabled: + if self.config.scan_layers: + return layer_output, None, cached_indexer_state + return layer_output, kv_cache, cached_indexer_state + + if self.config.scan_layers: + return layer_output, None + return layer_output, kv_cache + + def self_attention_with_norm_op( + self, + inputs, + decoder_segment_ids, + decoder_positions, + deterministic, + model_mode, + previous_chunk=None, + slot: None | int = None, + cached_indexer_state=None, + layer_idx=None, + ): + """Self-attention with normalization and IndexShare caching.""" + if self.is_mhc_enabled: + intermediate_inputs, _ = self.mhc_attention( + self.pre_attention_norm_op, + self.self_attention, + x=inputs, + mhc_type=HyperConnectionType.ATTENTION, + decoder_segment_ids=decoder_segment_ids, + inputs_positions=decoder_positions, + deterministic=deterministic, + model_mode=model_mode, + out_sharding=self.out_sharding, + previous_chunk=previous_chunk, + slot=slot, + cached_indexer_state=cached_indexer_state, + ) + new_indexer_state = None + else: + lnx = self.pre_attention_norm_op(inputs) + attention_lnx, new_indexer_state = self.attention_op( + lnx, + decoder_segment_ids, + decoder_positions, + deterministic, + model_mode, + previous_chunk, + slot, + cached_indexer_state=cached_indexer_state, + layer_idx=layer_idx, + ) + intermediate_inputs = inputs + attention_lnx + # Normalization + hidden_states = self.post_attention_norm_op(intermediate_inputs) + return hidden_states, intermediate_inputs, new_indexer_state + + +class GLMDenseLayer(GLMGenericLayer): + """GLM dense layer with Multi-Head Latent Attention.""" + + def __init__( + self, + config: Config, + model_mode: str, + mesh: Mesh, + rngs: nnx.Rngs, + quant: Optional[quantizations.AqtQuantization] = None, + layer_idx: int = -1, + ) -> None: + super().__init__(config, model_mode, mesh, rngs, quant, layer_idx) + self.mlp = linears.MlpBlock( + in_features=self.dummy_inputs_shape[-1], + intermediate_dim=self.config.mlp_dim, + activations=self.config.mlp_activations, + intermediate_dropout_rate=self.config.dropout_rate, + dtype=self.config.dtype, + weight_dtype=self.config.weight_dtype, + config=self.config, + quant=quant, + model_mode=model_mode, + mesh=mesh, + rngs=self.rngs, + ) + + def mlp_op(self, x, deterministic, *args, **kwargs): + mlp = self.mlp(x, deterministic, intermediate_sharding=self.mlp_intermediate_sharding, out_sharding=self.out_sharding) + return self.with_logical_constraint(mlp) + + def __call__( + self, + inputs, + decoder_segment_ids, + decoder_positions, + deterministic, + model_mode, + previous_chunk=None, + slot: None | int = None, + kv_cache=None, + attention_metadata=None, + decoder_input_tokens=None, + cached_indexer_state=None, + layer_idx=None, + **kwargs, + ): + if isinstance(inputs, tuple): + inputs = inputs[0] + x = self.with_logical_constraint(inputs) + x = checkpoint_name(x, "decoder_layer_input") + + if self.is_engram_enabled: + engram_output = self.engram_op(x, decoder_input_tokens) + x = x + engram_output + + hidden_states, intermediate_inputs, new_indexer_state = self.self_attention_with_norm_op( + x, + decoder_segment_ids, + decoder_positions, + deterministic, + model_mode, + previous_chunk, + slot, + cached_indexer_state=cached_indexer_state, + layer_idx=layer_idx, + ) + + if self.is_mhc_enabled: + layer_output, _ = self.mhc_mlp( + self.post_attention_norm_op, + self.mlp, + x=intermediate_inputs, + mhc_type=HyperConnectionType.MLP_DENSE, + deterministic=deterministic, + ) + else: + mlp_lnx = self.mlp_op(hidden_states, deterministic) + layer_output = mlp_lnx + intermediate_inputs + layer_output = self.dropout_op(layer_output, deterministic=deterministic) + + return self.post_process(layer_output, None, None, kv_cache, new_indexer_state) + + +class GLMMoELayer(GLMGenericLayer): + """GLM MoE layer with Multi-Head Latent Attention and IndexShare support.""" + + def __init__( + self, + config: Config, + model_mode: str, + mesh: Mesh, + rngs: nnx.Rngs, + quant: Optional[quantizations.AqtQuantization] = None, + layer_idx: int = -1, + ) -> None: + super().__init__(config, model_mode, mesh, rngs, quant, layer_idx) + self.DeepSeekMoeBlock_0 = moe.RoutedAndSharedMoE( + config=self.config, + mesh=mesh, + kernel_init=initializers.nd_dense_init(self.config.dense_init_scale, "fan_in", "truncated_normal"), + kernel_axes=("embed", None), + dtype=self.config.dtype, + weight_dtype=self.config.weight_dtype, + quant=quant, + rngs=self.rngs, + ) + + def mlp_op(self, x, deterministic, *args, **kwargs): + mlp_lnx, load_balance_loss, moe_bias_updates = self.DeepSeekMoeBlock_0( + x, intermediate_sharding=self.mlp_intermediate_sharding, out_sharding=self.out_sharding + ) + return self.with_logical_constraint(mlp_lnx), load_balance_loss, moe_bias_updates + + def __call__( + self, + inputs, + decoder_segment_ids, + decoder_positions, + deterministic, + model_mode, + previous_chunk=None, + slot: None | int = None, + kv_cache=None, + attention_metadata=None, + decoder_input_tokens=None, + cached_indexer_state=None, + layer_idx=None, + **kwargs, + ): + if isinstance(inputs, tuple): + inputs = inputs[0] + + x = self.with_logical_constraint(inputs) + x = checkpoint_name(x, "decoder_layer_input") + + if self.is_engram_enabled: + engram_output = self.engram_op(x, decoder_input_tokens) + x = x + engram_output + + hidden_states, intermediate_inputs, new_indexer_state = self.self_attention_with_norm_op( + x, + decoder_segment_ids, + decoder_positions, + deterministic, + model_mode, + previous_chunk, + slot, + cached_indexer_state=cached_indexer_state, + layer_idx=layer_idx, + ) + + if self.is_mhc_enabled: + layer_output, metadata = self.mhc_mlp( + self.post_attention_norm_op, + self.DeepSeekMoeBlock_0, + x=intermediate_inputs, + mhc_type=HyperConnectionType.MLP_MOE, + ) + load_balance_loss = metadata["load_balance_loss"] + moe_bias_updates = metadata["moe_bias_updates"] + else: + mlp_lnx, load_balance_loss, moe_bias_updates = self.mlp_op(hidden_states, deterministic) + layer_output = mlp_lnx + intermediate_inputs + layer_output = self.dropout_op(layer_output, deterministic=deterministic) + + return self.post_process(layer_output, load_balance_loss, moe_bias_updates, kv_cache, new_indexer_state) + + +GLMDenseLayerToLinen = nnx_wrappers.to_linen_class( + GLMDenseLayer, + base_metadata_fn=initializers.variable_to_logically_partitioned, +) + +GLMMoELayerToLinen = nnx_wrappers.to_linen_class( + GLMMoELayer, + base_metadata_fn=initializers.variable_to_logically_partitioned, +) diff --git a/src/maxtext/trainers/pre_train/train.py b/src/maxtext/trainers/pre_train/train.py index 274d9acc0a..e95cdb4920 100644 --- a/src/maxtext/trainers/pre_train/train.py +++ b/src/maxtext/trainers/pre_train/train.py @@ -850,6 +850,8 @@ def training_loop_iteration( if shard_optimizer_over_data and isinstance(model, nn.Module): state = sharding.maybe_shard_with_name(state, state_mesh_shardings, shard_mode) state, metrics = p_train_step(state, example_batch, *step_rng_args) + if prof.is_active(step): + jax.tree_util.tree_map(lambda x: x.block_until_ready() if hasattr(x, "block_until_ready") else x, metrics) step_time_delta = datetime.datetime.now() - last_step_completion last_step_completion = datetime.datetime.now() diff --git a/src/maxtext/utils/elastic_utils.py b/src/maxtext/utils/elastic_utils.py index 9d0d62fc28..fa6a2e1296 100644 --- a/src/maxtext/utils/elastic_utils.py +++ b/src/maxtext/utils/elastic_utils.py @@ -17,15 +17,22 @@ from collections import Counter import functools from types import SimpleNamespace +from typing import Any import jax from maxtext.utils import gcs_utils from maxtext.utils import max_logging -import pathwaysutils -from pathwaysutils.elastic import elastic -from pathwaysutils.elastic import manager -elastic_manager: manager.Manager | None = None +try: + import pathwaysutils + from pathwaysutils.elastic import elastic + from pathwaysutils.elastic import manager +except (ImportError, ModuleNotFoundError, AttributeError): + pathwaysutils = None + elastic = None + manager = None + +elastic_manager: Any | None = None pending_reinit_recorder = None pending_elastic_event_type = None @@ -35,6 +42,7 @@ def record_slice_state(recorder, active_slices_override: int | None = None) -> N if ( recorder is None or not hasattr(recorder, "record_elastic_slice_counts") + or pathwaysutils is None or not pathwaysutils.is_pathways_backend_used() or elastic_manager is None ): @@ -88,7 +96,7 @@ def record_elastic_reinit_end() -> None: def elastic_enabled(config) -> bool: """Returns whether elastic mode is enabled.""" - return pathwaysutils.is_pathways_backend_used() and config.elastic_enabled + return pathwaysutils is not None and pathwaysutils.is_pathways_backend_used() and config.elastic_enabled def elastic_snapshot(config) -> bool: diff --git a/src/maxtext/utils/globals.py b/src/maxtext/utils/globals.py index 30f6e65124..8cae6e65ed 100644 --- a/src/maxtext/utils/globals.py +++ b/src/maxtext/utils/globals.py @@ -91,6 +91,8 @@ "olmo3-7b": "allenai/Olmo-3-7B-Instruct", "olmo3-7b-pt": "allenai/Olmo-3-1025-7B", "olmo3-32b": "allenai/Olmo-3-32B-Think", + "glm5.1-744b": "zai-org/GLM-5.1", + "glm5.2-744b": "zai-org/GLM-5.2", # "default" is not HF model, but adding to to avoid confusing warning about tokenizer_path "default": os.path.join(MAXTEXT_ASSETS_ROOT, "tokenizers/tokenizer.llama2"), } diff --git a/src/maxtext/utils/index_share_utils.py b/src/maxtext/utils/index_share_utils.py new file mode 100644 index 0000000000..2801f07b92 --- /dev/null +++ b/src/maxtext/utils/index_share_utils.py @@ -0,0 +1,106 @@ +# 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. + +"""Utilities for GLM-5.2 Cross-Layer IndexCache (IndexShare). + +References: + GLM-5.2 / DSA IndexShare: Exploiting cross-layer token selection stability + to reduce lightning indexer compute by 50% to 75%. +""" + +from typing import Sequence + + +def parse_index_share_pattern(pattern: str | Sequence[str], num_layers: int) -> tuple[str, ...]: + """Parses and validates the IndexShare pattern string. + + Args: + pattern: Pattern string (e.g. "FSSS", "F,S,S,S", "FSSSFSSSFSSS...") or list of roles. + num_layers: Total number of decoder layers in the model. + + Returns: + A tuple of 'F' (Full layer) and 'S' (Shared layer) strings of length `num_layers`. + + Raises: + ValueError: If pattern is empty, contains invalid characters, or layer 0 is not 'F'. + """ + if isinstance(pattern, str): + # Normalize commas/spaces/case + clean_pattern = pattern.replace(",", "").replace(" ", "").upper() + else: + clean_pattern = "".join(str(x).strip().upper() for x in pattern) + + if not clean_pattern: + raise ValueError("index_share_pattern cannot be empty.") + + invalid_chars = set(clean_pattern) - {"F", "S"} + if invalid_chars: + raise ValueError( + f"Invalid characters in index_share_pattern: {invalid_chars}. Only 'F' (Full) and 'S' (Shared) are allowed." + ) + + if clean_pattern[0] != "F": + raise ValueError(f"First layer (Layer 0) must always be 'F' (Full layer), but got '{clean_pattern[0]}'.") + + # If pattern is shorter than num_layers, repeat it periodically to fill num_layers + if len(clean_pattern) < num_layers: + repeats = (num_layers + len(clean_pattern) - 1) // len(clean_pattern) + full_pattern = (clean_pattern * repeats)[:num_layers] + elif len(clean_pattern) > num_layers: + full_pattern = clean_pattern[:num_layers] + else: + full_pattern = clean_pattern + + return tuple(full_pattern) + + +def get_donor_layer_indices(pattern_tuple: tuple[str, ...]) -> tuple[int, ...]: + """For each layer, returns the index of its donor Full (F) layer. + + f(l) = max{ j <= l : pattern[j] == 'F' } + """ + donor_indices = [] + current_f = 0 + for idx, role in enumerate(pattern_tuple): + if role == "F": + current_f = idx + donor_indices.append(current_f) + return tuple(donor_indices) + + +def get_donor_layer_idx(layer_idx: int, pattern_tuple: tuple[str, ...]) -> int: + """Returns the donor Full (F) layer index for a specific layer.""" + return get_donor_layer_indices(pattern_tuple)[layer_idx] + + +def get_served_group_sizes(pattern_tuple: tuple[str, ...]) -> tuple[int, ...]: + """For each layer, returns the group size |Served(f(l))| of its donor F-layer. + + This is used to normalize the multi-layer distillation loss: + L_multi_I = 1 / |Served(l)| * sum_{j in Served(l)} KL(p^(j) || q^(l)) + """ + donor_indices = get_donor_layer_indices(pattern_tuple) + # Count how many layers each donor F-layer serves + counts: dict[int, int] = {} + for d in donor_indices: + counts[d] = counts.get(d, 0) + 1 + + return tuple(counts[d] for d in donor_indices) + + +def is_shared_layer(layer_idx: int, pattern_tuple: tuple[str, ...]) -> bool: + """Returns True if the given layer index is a Shared (S) layer.""" + if layer_idx < 0 or layer_idx >= len(pattern_tuple): + return False + return pattern_tuple[layer_idx] == "S" diff --git a/src/maxtext/utils/maxtext_utils.py b/src/maxtext/utils/maxtext_utils.py index 76a9c842e2..e579d7a876 100644 --- a/src/maxtext/utils/maxtext_utils.py +++ b/src/maxtext/utils/maxtext_utils.py @@ -745,7 +745,7 @@ def calculate_routed_and_shared_ffn_tflops_per_device(config): def get_dense_moe_layers(config): """Helper function to calculate number of dense and moe layers""" - if config.decoder_block == DecoderBlockType.DEEPSEEK: + if config.decoder_block in (DecoderBlockType.DEEPSEEK, DecoderBlockType.GLM5): num_dense_layers = config.first_num_dense_layers num_moe_layers = config.num_decoder_layers - config.first_num_dense_layers return num_dense_layers, num_moe_layers @@ -1154,6 +1154,7 @@ def calculate_tflops_training_per_device(config, log=True): # calculation based on dropless implementation if config.decoder_block in ( DecoderBlockType.DEEPSEEK, + DecoderBlockType.GLM5, DecoderBlockType.LLAMA4, DecoderBlockType.QWEN3_NEXT, DecoderBlockType.QWEN3_5, @@ -1253,7 +1254,7 @@ def calculate_tflops_training_per_device(config, log=True): attention_tflops, learnable_weight_tflops = calculate_deepseek4_tflops_training_per_device( config, total_ffn_flops_all_layers, embedding_flops ) - elif config.decoder_block == DecoderBlockType.DEEPSEEK: + elif config.decoder_block in (DecoderBlockType.DEEPSEEK, DecoderBlockType.GLM5): learnable_weight_tflops = ( (total_ffn_flops_all_layers + (qkv_flops + projection_flops) * config.num_decoder_layers + embedding_flops) * 3 @@ -1320,6 +1321,7 @@ def calculate_tflops_training_per_device(config, log=True): gate_flops = 2 * config.per_device_batch_size * config.max_target_length * config.emb_dim * config.num_experts if config.decoder_block in ( DecoderBlockType.DEEPSEEK, + DecoderBlockType.GLM5, DecoderBlockType.LLAMA4, DecoderBlockType.QWEN3_NEXT, DecoderBlockType.GEMMA4, diff --git a/tests/end_to_end/tpu/glm5/glm5.2-744b/1_test_glm5.sh b/tests/end_to_end/tpu/glm5/glm5.2-744b/1_test_glm5.sh new file mode 100644 index 0000000000..5037b1558f --- /dev/null +++ b/tests/end_to_end/tpu/glm5/glm5.2-744b/1_test_glm5.sh @@ -0,0 +1,54 @@ +#!/bin/bash + +# This file is documentation for how to get started with GLM-5.2 (Cross-Layer IndexShare). + +# This file runs Step 1 on CPU. +# 1. Convert the HuggingFace checkpoint (bf16) to MaxText-compatible checkpoint (bf16): +# Scanned format is better for training; unscanned format is better for decoding. +# 2. Run logit check, pre-training, fine-tuning, and decoding. + +set -ex + +export MODEL_NAME='glm5.2-744b' +export TOKENIZER_PATH='zai-org/GLM-5.2' + +# Installing torch for checkpoint conversion and forward_pass_logit_checker.py +python3 -m pip install torch --index-url https://download.pytorch.org/whl/cpu + +if [ -z "${BASE_OUTPUT_PATH}" ]; then + export BASE_OUTPUT_PATH=gs://runner-maxtext-logs/$(date +%Y-%m-%d-%H-%M) + echo "BASE_OUTPUT_PATH is not set" +fi +BASE_OUTPUT_PATH=${BASE_OUTPUT_PATH%/} +echo using BASE_OUTPUT_PATH = ${BASE_OUTPUT_PATH} + +# Step 1: Checkpoint conversion +# HF checkpoint: https://huggingface.co/zai-org/GLM-5.2 +BF16_HF_PATH=${BF16_HF_PATH:-gs://maxtext-glm5-europe-west4/glm5.2_raw} +if [ -z "${BF16_LOCAL_PATH}" ] && [ ! -d "/home/rishabhbaghel_google_com/glm5.2_raw" ]; then + export BF16_LOCAL_PATH=/tmp/glm5.2_raw + gcloud storage cp -r ${BF16_HF_PATH} /tmp || true +fi +BF16_LOCAL_PATH=${BF16_LOCAL_PATH:-/home/rishabhbaghel_google_com/glm5.2_raw} + +# scanned +python3 -m maxtext.checkpoint_conversion.to_maxtext src/maxtext/configs/base.yml \ + model_name=${MODEL_NAME} scan_layers=true \ + base_output_directory=${BASE_OUTPUT_PATH}/scanned hf_access_token=$HF_TOKEN \ + hardware=cpu skip_jax_distributed_system=True \ + checkpoint_storage_concurrent_gb=1024 \ + --hf_model_path=$BF16_LOCAL_PATH \ + --lazy_load_tensors=False \ + --eager_load_method=safetensors \ + --save_dtype=bfloat16 + +# unscanned +python3 -m maxtext.checkpoint_conversion.to_maxtext src/maxtext/configs/base.yml \ + model_name=${MODEL_NAME} scan_layers=false \ + base_output_directory=${BASE_OUTPUT_PATH}/unscanned hf_access_token=$HF_TOKEN \ + hardware=cpu skip_jax_distributed_system=True \ + checkpoint_storage_concurrent_gb=1024 \ + --hf_model_path=$BF16_LOCAL_PATH \ + --lazy_load_tensors=False \ + --eager_load_method=safetensors \ + --save_dtype=bfloat16 diff --git a/tests/end_to_end/tpu/glm5/glm5.2-744b/2_test_glm5.sh b/tests/end_to_end/tpu/glm5/glm5.2-744b/2_test_glm5.sh new file mode 100644 index 0000000000..67c9cb9a9d --- /dev/null +++ b/tests/end_to_end/tpu/glm5/glm5.2-744b/2_test_glm5.sh @@ -0,0 +1,80 @@ +#!/bin/bash + +# This file runs Step 2 on TPU cluster for GLM-5.2 (Cross-Layer IndexShare). +# 1. Forward pass logit check against golden logits. +# 2. High-throughput distributed pre-training with IndexShare. +# 3. Decoding & sanity prompt generation. + +set -ex + +export MODEL_NAME='glm5.2-744b' +export TOKENIZER_PATH='zai-org/GLM-5.2' + +# Installing torch CPU for tokenizer / evaluation helpers +python3 -m pip install torch --index-url https://download.pytorch.org/whl/cpu + +if [ -z "${BASE_OUTPUT_PATH}" ]; then + export BASE_OUTPUT_PATH=gs://runner-maxtext-logs/$(date +%Y-%m-%d-%H-%M) + echo "BASE_OUTPUT_PATH is not set" +fi +BASE_OUTPUT_PATH=${BASE_OUTPUT_PATH%/} +echo using BASE_OUTPUT_PATH = ${BASE_OUTPUT_PATH} + +SCANNED_CKPT_PATH=${SCANNED_CKPT_PATH:-gs://maxtext-glm5-europe-west4/maxtext-glm-5.2-bf16-converted-final-78l/0/items} +export DATASET_PATH=gs://maxtext-dataset + +# 1. Distributed Pre-Training Benchmark with GLM-5.2 Cross-Layer IndexShare (64 TPU cores, EP=4, FSDP=16) +python3 -m maxtext.trainers.pre_train.train src/maxtext/configs/base.yml \ + base_output_directory=${BASE_OUTPUT_PATH} \ + run_name=pretrain_glm52 \ + model_name=${MODEL_NAME} \ + override_model_config=true \ + dataset_type=synthetic \ + tokenizer_type=huggingface \ + tokenizer_path=${TOKENIZER_PATH} \ + per_device_batch_size=1 \ + max_target_length=2048 \ + indexer_topk=1024 \ + use_indexer=true \ + use_index_share=true \ + index_share_pattern="FSSS" \ + prune_shared_indexers=true \ + dcn_pipeline_parallelism=1 \ + dcn_data_parallelism=-1 \ + ici_pipeline_parallelism=1 \ + ici_fsdp_transpose_parallelism=1 \ + ici_fsdp_parallelism=16 \ + ici_expert_parallelism=4 \ + allow_split_physical_axes=true \ + use_iota_embed=true \ + remat_policy=custom \ + decoder_layer_input=offload \ + opt_type=sgd \ + enable_checkpointing=false \ + steps=10 + +# 2. Forward Logit & Generation Test with GLM-5.2 Cross-Layer IndexShare +python3 -m maxtext.inference.decode src/maxtext/configs/base.yml \ + base_output_directory=${BASE_OUTPUT_PATH} \ + run_name=decode_glm52 \ + model_name=${MODEL_NAME} \ + tokenizer_type=huggingface \ + tokenizer_path=${TOKENIZER_PATH} \ + load_parameters_path=${SCANNED_CKPT_PATH} \ + scan_layers=true \ + attention=dot_product \ + sparse_matmul=false \ + dtype=bfloat16 \ + weight_dtype=bfloat16 \ + per_device_batch_size=1 \ + max_prefill_predict_length=64 \ + max_target_length=128 \ + ici_fsdp_parallelism=16 \ + ici_expert_parallelism=4 \ + checkpoint_storage_concurrent_gb=1024 \ + use_indexer=true \ + use_index_share=true \ + index_share_pattern="FSSS" \ + prune_shared_indexers=true \ + prompt="The capital of France is" + diff --git a/tests/unit/glm52_indexshare_test.py b/tests/unit/glm52_indexshare_test.py new file mode 100644 index 0000000000..931e746d52 --- /dev/null +++ b/tests/unit/glm52_indexshare_test.py @@ -0,0 +1,72 @@ +# 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. + +"""Unit tests for GLM-5.2 Training-Aware IndexShare (Cross-Layer IndexCache).""" + +import unittest +from maxtext.utils import index_share_utils + + +class GLM52IndexSharePatternTest(unittest.TestCase): + """Tests for IndexShare pattern utilities.""" + + def test_pattern_expansion_and_validation(self): + # Test periodic expansion for 78 layers (GLM-5.1/5.2 default) + pattern = index_share_utils.parse_index_share_pattern("FSSS", 78) + self.assertEqual(len(pattern), 78) + self.assertEqual(pattern[0], "F") + self.assertEqual(pattern[1], "S") + self.assertEqual(pattern[2], "S") + self.assertEqual(pattern[3], "S") + self.assertEqual(pattern[4], "F") + + # Count F and S layers (1/4 retention) + num_f = sum(1 for p in pattern if p == "F") + num_s = sum(1 for p in pattern if p == "S") + self.assertEqual(num_f, 20) # ceil(78/4) + self.assertEqual(num_s, 58) + + def test_donor_mapping(self): + pattern = index_share_utils.parse_index_share_pattern("FSSS", 8) + donors = index_share_utils.get_donor_layer_indices(pattern) + self.assertEqual(donors, (0, 0, 0, 0, 4, 4, 4, 4)) + + def test_group_sizes(self): + pattern = index_share_utils.parse_index_share_pattern("FSSS", 8) + sizes = index_share_utils.get_served_group_sizes(pattern) + self.assertEqual(sizes, (4, 4, 4, 4, 4, 4, 4, 4)) + + def test_invalid_pattern_raises(self): + with self.assertRaises(ValueError): + index_share_utils.parse_index_share_pattern("SFFF", 4) + with self.assertRaises(ValueError): + index_share_utils.parse_index_share_pattern("FABCS", 5) + with self.assertRaises(ValueError): + index_share_utils.parse_index_share_pattern("", 4) + + def test_checkpoint_donor_resolution(self): + pattern = index_share_utils.parse_index_share_pattern("FSSS", 12) + # Layers 0..3 share with Layer 0 + for l in range(4): + self.assertEqual(index_share_utils.get_donor_layer_idx(l, pattern), 0) + # Layers 4..7 share with Layer 4 + for l in range(4, 8): + self.assertEqual(index_share_utils.get_donor_layer_idx(l, pattern), 4) + # Layers 8..11 share with Layer 8 + for l in range(8, 12): + self.assertEqual(index_share_utils.get_donor_layer_idx(l, pattern), 8) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/unit/index_share_utils_test.py b/tests/unit/index_share_utils_test.py new file mode 100644 index 0000000000..04c3ccea2b --- /dev/null +++ b/tests/unit/index_share_utils_test.py @@ -0,0 +1,68 @@ +# 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. + +"""Unit tests for IndexShare pattern utilities.""" + +import unittest +from maxtext.utils import index_share_utils + + +class IndexShareUtilsTest(unittest.TestCase): + + def test_parse_index_share_pattern_periodic(self): + pattern = index_share_utils.parse_index_share_pattern("FSSS", 10) + self.assertEqual(len(pattern), 10) + self.assertEqual(pattern, ("F", "S", "S", "S", "F", "S", "S", "S", "F", "S")) + + def test_parse_index_share_pattern_with_commas_and_spaces(self): + pattern = index_share_utils.parse_index_share_pattern("f, s, s, s", 8) + self.assertEqual(pattern, ("F", "S", "S", "S", "F", "S", "S", "S")) + + def test_parse_index_share_pattern_exact(self): + pattern = index_share_utils.parse_index_share_pattern("FSFSS", 5) + self.assertEqual(pattern, ("F", "S", "F", "S", "S")) + + def test_invalid_first_layer(self): + with self.assertRaises(ValueError) as ctx: + index_share_utils.parse_index_share_pattern("SFFF", 4) + self.assertIn("First layer (Layer 0) must always be 'F'", str(ctx.exception)) + + def test_invalid_characters(self): + with self.assertRaises(ValueError) as ctx: + index_share_utils.parse_index_share_pattern("FXSS", 4) + self.assertIn("Invalid characters", str(ctx.exception)) + + def test_donor_indices(self): + pattern = ("F", "S", "S", "S", "F", "S", "S") + donors = index_share_utils.get_donor_layer_indices(pattern) + self.assertEqual(donors, (0, 0, 0, 0, 4, 4, 4)) + + def test_group_sizes(self): + pattern = ("F", "S", "S", "S", "F", "S", "S") + sizes = index_share_utils.get_served_group_sizes(pattern) + # Layer 0 serves 4 layers (0, 1, 2, 3) -> size 4 + # Layer 4 serves 3 layers (4, 5, 6) -> size 3 + self.assertEqual(sizes, (4, 4, 4, 4, 3, 3, 3)) + + def test_is_shared_layer(self): + pattern = ("F", "S", "S", "F") + self.assertFalse(index_share_utils.is_shared_layer(0, pattern)) + self.assertTrue(index_share_utils.is_shared_layer(1, pattern)) + self.assertTrue(index_share_utils.is_shared_layer(2, pattern)) + self.assertFalse(index_share_utils.is_shared_layer(3, pattern)) + self.assertFalse(index_share_utils.is_shared_layer(4, pattern)) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/unit/nnx_decoders_test.py b/tests/unit/nnx_decoders_test.py index c2832f4ce9..8e4040eca4 100644 --- a/tests/unit/nnx_decoders_test.py +++ b/tests/unit/nnx_decoders_test.py @@ -1293,7 +1293,6 @@ def __call__(self, x, **kwargs): maxtext_utils_nnx.nnx_add_and_sync_scan_axis = mock_add_scan_axis try: - # Use a custom metadata_axis_name custom_axis_name = "custom_scanned_blocks" # pylint: disable=protected-access _, _, _ = decoder._apply_layers_sequentially( @@ -1302,11 +1301,9 @@ def __call__(self, x, **kwargs): length=2, metadata_axis_name=custom_axis_name, ) - - # Verify that the custom axis name was indeed passed down found_custom_name = False for call_args in mock_add_scan_axis.call_args_list: - if call_args[0][1] == custom_axis_name: + if len(call_args[0]) > 1 and call_args[0][1] == custom_axis_name: found_custom_name = True break diff --git a/tests/unit/pallas_mosaic_tpu_v2_kernel_test.py b/tests/unit/pallas_mosaic_tpu_v2_kernel_test.py index 9c16084e24..b200aeb458 100644 --- a/tests/unit/pallas_mosaic_tpu_v2_kernel_test.py +++ b/tests/unit/pallas_mosaic_tpu_v2_kernel_test.py @@ -19,7 +19,8 @@ from absl.testing import absltest from absl.testing import parameterized -import chex + +chex = pytest.importorskip("chex", reason="chex not installed") import jax from jax.experimental import pallas as pl from jax.experimental.pallas import tpu as pltpu diff --git a/tests/utils/forward_pass_logit_checker.py b/tests/utils/forward_pass_logit_checker.py index a51b23980f..70d5d678ea 100644 --- a/tests/utils/forward_pass_logit_checker.py +++ b/tests/utils/forward_pass_logit_checker.py @@ -271,6 +271,7 @@ def get_data(golden_data_point, config): model_prefix = config.model_name.split("-")[0] max_logging.log(f"config.global_batch_size_to_train_on={config.global_batch_size_to_train_on}") + model_prefix = config.model_name.split("-")[0] if config.use_multimodal: assert "pixel_values" in golden_data_point, "no image found in golden data while use_multimodal=True" pixel_values = np.asarray(golden_data_point["pixel_values"], dtype=np.float32)