diff --git a/src/maxtext/configs/base.yml b/src/maxtext/configs/base.yml index 9c834fb7d4..e74b3db388 100644 --- a/src/maxtext/configs/base.yml +++ b/src/maxtext/configs/base.yml @@ -356,6 +356,7 @@ moe_mlpwo: 'remat' query_proj: 'remat' key_proj: 'remat' value_proj: 'remat' +kv_proj: 'remat' qkv_proj: 'remat' out_proj: 'remat' query_wa_proj: 'remat' @@ -1133,7 +1134,7 @@ enable_prefix_caching: false prefix_caching_hbm_byte: 10_000_000_000 # 10 GB prefix_caching_dram_byte: 100_000_000_000 # 100 GB -# This is a temporary flag that will be removed soon after the fix lands in TE +# This is a temporary flag that will be removed soon after the fix lands in Transformer Engine enable_padding_causal_mask: true # Llama4-specific diff --git a/src/maxtext/configs/types.py b/src/maxtext/configs/types.py index 3dffe5e7ac..b6cb69918f 100644 --- a/src/maxtext/configs/types.py +++ b/src/maxtext/configs/types.py @@ -518,7 +518,10 @@ class ModelArchitecture(BaseModel): True, description="Whether to apply scale on query and key normalizations (default True).", ) - v_norm_with_scale: bool = Field(True, description="Whether to apply scale on value normalization (default True).") + v_norm_with_scale: bool = Field( + True, + description="Whether to apply scale on value normalization (default True).", + ) class MTP(BaseModel): @@ -624,7 +627,7 @@ class Attention(BaseModel): ), ) ragged_block_size: int = Field(256, description="Block size for ragged attention.") - enable_padding_causal_mask: bool = Field(True, description="Temporary flag for TE padding.") + enable_padding_causal_mask: bool = Field(True, description="Temporary flag for Transformer Engine padding.") use_tokamax_splash: bool = Field(False, description="Whether to use tokamax splash attention.") use_jax_splash: bool = Field(False, description="Whether to use jax splash attention.") force_q_layout: bool = Field(False, description="Force the Q layout") @@ -656,7 +659,10 @@ class CompressedAttention(BaseModel): o_lora_rank: NonNegativeInt = Field(0, description="Output LoRA rank for Compressed Attention.") o_groups: NonNegativeInt = Field(0, description="Output groups for Compressed Attention.") - compress_ratios: list[int] = Field(default_factory=list, description="Per-layer compression ratios (0, 4, 128, etc).") + compress_ratios: list[int] = Field( + default_factory=list, + description="Per-layer compression ratios (0, 4, 128, etc).", + ) compressed_rope_max_timescale: int = Field( 160000, description="If positive, used for Compressed Sparse/Heavy Attention." ) @@ -757,14 +763,18 @@ class MoEGeneral(BaseModel): num_experts: PositiveInt = Field(1, description="The total number of experts in each MoE layer.") num_experts_per_tok: PositiveInt = Field(1, description="The number of experts to route each token to.") capacity_factor: float = Field(-1.0, description="Expert capacity factor. If < 0, no token dropping.") - ragged_buffer_factor: float = Field(-1.0, description="Ragged buffer factor. If < 0, ragged buffer is worst case size.") + ragged_buffer_factor: float = Field( + -1.0, + description="Ragged buffer factor. If < 0, ragged buffer is worst case size.", + ) moe_expert_input_dim: int = Field( -1, description="Dimension of tokens entering the MoE layer. If < 0, defaults to emb_dim.", ) base_moe_mlp_dim: int = Field(-1, description="Intermediate dimension at MoE layer.") padded_base_moe_mlp_dim: Optional[int] = Field( - None, description="Padded intermediate dimension at MoE layer for efficient GMM_v2 kernel execution." + None, + description="Padded intermediate dimension at MoE layer for efficient GMM_v2 kernel execution.", ) load_balance_loss_weight: NonNegativeFloat = Field(0.0, description="Weight for the load balancing auxiliary loss.") use_custom_sort_vjp: bool = Field( @@ -783,7 +793,8 @@ class MoEGeneral(BaseModel): ), ) use_ragged_sort: bool = Field( - False, description="Whether to use ragged kernel for sorting, improve performance when EP is enabled." + False, + description="Whether to use ragged kernel for sorting, improve performance when EP is enabled.", ) use_gather_mosaic_kernel: bool = Field( False, @@ -919,7 +930,8 @@ class DeepSeekMoE(BaseModel): ) n_routing_groups: int = Field(-1, description="Number of groups for routing, disabled by default.") first_num_hash_layers: int = Field( - 0, description="Number of hash routing layers, used in DeepSeek V4 (0 means disabled)." + 0, + description="Number of hash routing layers, used in DeepSeek V4 (0 means disabled).", ) topk_routing_group: int = Field(-1, description="Number of top groups to route inputs to.") use_batch_split_schedule: bool = Field( @@ -993,7 +1005,8 @@ class HardwareAndMesh(BaseModel): ) custom_mesh: str = Field("", description="Available options: ['hybrid_ring_64x4', 'hybrid_ring_32x8']") custom_mesh_and_rule: CustomRule = Field( - CustomRule.DEFAULT, description="Customized mesh and logical rules for granularity." + CustomRule.DEFAULT, + description="Customized mesh and logical rules for granularity.", ) allow_split_physical_axes: bool = Field(False, description="Allow splitting physical axes for device mesh creation.") enable_nnx: bool = Field(True, description="Whether to use NNX for model definition.") @@ -1002,7 +1015,8 @@ class HardwareAndMesh(BaseModel): pure_nnx_decoder: bool = Field(True, description="Whether to enable pure NNX decoder.") pure_nnx: bool = Field(True, description="Whether to enable pure NNX mode.") remove_size_one_mesh_axis_from_type: bool = Field( - True, description="Whether to remove size one mesh axis from type through jax.config." + True, + description="Whether to remove size one mesh axis from type through jax.config.", ) @@ -1028,7 +1042,10 @@ class LayoutAndSharding(BaseModel): "with auto sharding, megablox kernel, and EP / FSDP parallelisms.", ) shard_optimizer_over_data: bool = Field(False, description="Enable ZeRO-1 optimizer sharding over the data axis.") - internal_compile: bool = Field(False, description="Use internal_compile to bypass open-source topology mappings.") + internal_compile: bool = Field( + False, + description="Use internal_compile to bypass open-source topology mappings.", + ) internal_compile_num_devices: int = Field(-1, description="Number of devices when using internal_compile.") compile_xla_flags: str = Field("", description="Compiler options for compilation only.") @@ -1073,7 +1090,8 @@ class PipelineParallelism(BaseModel): """Configuration for pipeline parallelism.""" pipeline_fsdp_ag_per_repeat: bool = Field( - False, description="Enable weight prefetching for circular pipeline parallelism." + False, + description="Enable weight prefetching for circular pipeline parallelism.", ) num_layers_per_pipeline_stage: int = Field(1, description="Number of layers to place on each pipeline stage.") num_pipeline_repeats: int = Field( @@ -1139,6 +1157,7 @@ class RematAndOffload(BaseModel): query_proj: RematLocation = Field(RematLocation.REMAT, description="Remat policy for the query projection.") key_proj: RematLocation = Field(RematLocation.REMAT, description="Remat policy for the key projection.") value_proj: RematLocation = Field(RematLocation.REMAT, description="Remat policy for the value projection.") + kv_proj: RematLocation = Field(RematLocation.REMAT, description="Remat policy for the unified KV projection.") query_wa_proj: RematLocation = Field( RematLocation.REMAT, description="Remat policy for the MLA query weighted attention projection.", @@ -1321,7 +1340,10 @@ class OlmoGrainDataset(BaseModel): ``data_shuffle_seed``); only OLMo-specific fields are listed here. """ - olmo_index_path: PathStr = Field("", description="Path or gs:// URI to the JSON index from build_olmo_npy_index.py.") + olmo_index_path: PathStr = Field( + "", + description="Path or gs:// URI to the JSON index from build_olmo_npy_index.py.", + ) olmo_path_remap_from: PathStr = Field( "", description="If set, rewrite index file paths starting with this prefix to olmo_path_remap_to.", @@ -1428,19 +1450,24 @@ class Distillation(BaseModel): distill_layer_indices: None | list = Field(None, description="Feature indices for feature loss.") distill_alpha_end: Optional[float] = Field(None, description="Target alpha at end of training. None keeps alpha fixed.") distill_alpha_schedule: Literal["constant", "linear", "cosine"] = Field( - "constant", description="Schedule type for alpha annealing ('constant', 'linear', or 'cosine')." + "constant", + description="Schedule type for alpha annealing ('constant', 'linear', or 'cosine').", ) distill_temperature_end: Optional[float] = Field( - None, description="Target temperature at end of training. None keeps temperature fixed." + None, + description="Target temperature at end of training. None keeps temperature fixed.", ) distill_temperature_schedule: Literal["constant", "linear", "cosine"] = Field( - "constant", description="Schedule type for temperature annealing ('constant', 'linear', or 'cosine')." + "constant", + description="Schedule type for temperature annealing ('constant', 'linear', or 'cosine').", ) distill_beta_end: Optional[float] = Field( - None, description="Target beta_feature at end of training. None keeps beta fixed." + None, + description="Target beta_feature at end of training. None keeps beta fixed.", ) distill_beta_schedule: Literal["constant", "linear", "cosine"] = Field( - "constant", description="Schedule type for beta annealing ('constant', 'linear', or 'cosine')." + "constant", + description="Schedule type for beta annealing ('constant', 'linear', or 'cosine').", ) # --- Learn to init related parameters -- @@ -1463,11 +1490,13 @@ class Distillation(BaseModel): ) attn_module_name: Optional[str] = Field( - None, description="Attention nnx module attribute name to augment with LTI logic" + None, + description="Attention nnx module attribute name to augment with LTI logic", ) lti_layer_indices: Optional[list[int]] = Field( - None, description="List of layer indices to apply LTI modifications. If None, applied to all layers." + None, + description="List of layer indices to apply LTI modifications. If None, applied to all layers.", ) # --------------------------------------- @@ -1532,11 +1561,13 @@ class DilocoParams(BaseModel): diloco_outer_lr: float = Field(0.3, description="learning rate for outer optimizer.") diloco_outer_momentum: float = Field(0.9, description="momentum for outer optimizer.") dcn_bandwidth_limit: str = Field( - "", description="Programmatic DCN egress bandwidth limit (e.g., '28gbit'). Empty means no limit." + "", + description="Programmatic DCN egress bandwidth limit (e.g., '28gbit'). Empty means no limit.", ) dcn_bandwidth_burst: str = Field("10mb", description="Burst size for Token Bucket Filter (TBF) traffic shaping.") dcn_bandwidth_latency: str = Field( - "50ms", description="Latency threshold for Token Bucket Filter (TBF) traffic shaping." + "50ms", + description="Latency threshold for Token Bucket Filter (TBF) traffic shaping.", ) dcn_bandwidth_interface: str = Field("eth0", description="Network interface to apply bandwidth limits on.") @@ -1829,7 +1860,8 @@ class Profiling(BaseModel): tpu_num_chips_to_profile_per_task: int = Field(1, description="Specifies the number of TPU chips to profile per task.") tpu_num_sparse_cores_to_trace: int = Field(2, description="Specifies the number of TPU chips to profile per task.") tpu_num_sparse_core_tiles_to_trace: int = Field( - 1, description="Specifies the number of tiles within each sparse core to trace on the TPU." + 1, + description="Specifies the number of tiles within each sparse core to trace on the TPU.", ) xprof_tpu_power_trace_level: XProfTPUPowerTraceMode = Field( XProfTPUPowerTraceMode.POWER_TRACE_NONE, @@ -2770,7 +2802,11 @@ def validate_and_set_hlo_dump_defaults(): ) for param_name, schedule, end_value in [ ("distill_alpha", self.distill_alpha_schedule, self.distill_alpha_end), - ("distill_temperature", self.distill_temperature_schedule, self.distill_temperature_end), + ( + "distill_temperature", + self.distill_temperature_schedule, + self.distill_temperature_end, + ), ("distill_beta", self.distill_beta_schedule, self.distill_beta_end), ]: if schedule != "constant" and end_value is None: @@ -2827,7 +2863,10 @@ def get_num_target_devices(): # Check for AQT deprecation warning if self.quantization and not self.use_qwix_quantization: - if self.quantization not in ("fp8", "nanoo_fp8") and not self.quantization.startswith("te_"): + if self.quantization not in ( + "fp8", + "nanoo_fp8", + ) and not self.quantization.startswith("te_"): logger.warning( "WARNING: AQT quantization is deprecated and will be removed in a future release. " "Please migrate to Qwix by setting use_qwix_quantization=True." @@ -2935,6 +2974,7 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de "query_proj", "key_proj", "value_proj", + "kv_proj", "query_wa_proj", "kv_wa_proj", "mla_kv", diff --git a/src/maxtext/layers/attention_compressed.py b/src/maxtext/layers/attention_compressed.py index 65869fb186..d8975f92a0 100644 --- a/src/maxtext/layers/attention_compressed.py +++ b/src/maxtext/layers/attention_compressed.py @@ -19,6 +19,7 @@ import jax import jax.numpy as jnp +from jax.ad_checkpoint import checkpoint_name from jax.sharding import Mesh from flax import nnx @@ -107,10 +108,18 @@ def csa_overlap_pooling( # Shift Ca forward by one window to align with the next Cb a_kv_shifted = jnp.concatenate( - [jnp.zeros((batch_size, 1, compress_rate, head_dim), dtype=a_kv.dtype), a_kv[:, :-1]], axis=1 + [ + jnp.zeros((batch_size, 1, compress_rate, head_dim), dtype=a_kv.dtype), + a_kv[:, :-1], + ], + axis=1, ) a_gate_shifted = jnp.concatenate( - [jnp.full((batch_size, 1, compress_rate, head_dim), -jnp.inf, dtype=a_gate.dtype), a_gate[:, :-1]], axis=1 + [ + jnp.full((batch_size, 1, compress_rate, head_dim), -jnp.inf, dtype=a_gate.dtype), + a_gate[:, :-1], + ], + axis=1, ) # Concatenate shifted Ca and unshifted Cb to form the final overlapping window @@ -239,7 +248,16 @@ def __init__( model_mode: The operational mode (e.g., "train", "prefill"). rngs: An optional Rngs instance for stochastic initializations or dropout. """ - super().__init__(config, compress_ratio, rotary_embedding, 1, kernel_init, quant, model_mode, rngs) + super().__init__( + config, + compress_ratio, + rotary_embedding, + 1, + kernel_init, + quant, + model_mode, + rngs, + ) def __call__( self, @@ -459,13 +477,19 @@ def __call__( compressed = self.rotary_emb(compressed, positions, unsqueeze_dim=None) else: # Return empty top-k selections when sequence is too short to form any windows - return jnp.zeros((batch_size, seq_len, min(self.index_topk, compressed_len)), dtype=jnp.int32) + return jnp.zeros( + (batch_size, seq_len, min(self.index_topk, compressed_len)), + dtype=jnp.int32, + ) # Broadcast the compressed KV representations across all indexer heads # -> [batch, 1, n_windows, index_head_dim] compressed_kv = jnp.expand_dims(compressed, axis=1) # -> [batch, index_n_heads, n_windows, index_head_dim] - compressed_kv = jnp.broadcast_to(compressed_kv, (batch_size, self.index_n_heads, compressed_len, self.index_head_dim)) + compressed_kv = jnp.broadcast_to( + compressed_kv, + (batch_size, self.index_n_heads, compressed_len, self.index_head_dim), + ) # Project the latent query to match the Indexer's dimensions # [batch, seq_len, index_n_heads * index_head_dim] -> [batch, seq_len, index_n_heads, index_head_dim] @@ -551,7 +575,16 @@ def __init__( model_mode: The operational mode (e.g., "train", "prefill"). rngs: An optional Rngs instance for stochastic initializations or dropout. """ - super().__init__(config, compress_ratio, rotary_embedding, 2, kernel_init, quant, model_mode, rngs) + super().__init__( + config, + compress_ratio, + rotary_embedding, + 2, + kernel_init, + quant, + model_mode, + rngs, + ) self.indexer = DeepseekV4Indexer( config=config, @@ -631,7 +664,11 @@ def __call__( compressed_mask = jnp.where(is_selected, 0.0, DEFAULT_MASK_VALUE).astype(self.dtype) else: - compressed_mask = jnp.full((batch_size, 1, seq_len, compressed_len), DEFAULT_MASK_VALUE, dtype=self.dtype) + compressed_mask = jnp.full( + (batch_size, 1, seq_len, compressed_len), + DEFAULT_MASK_VALUE, + dtype=self.dtype, + ) return compressed_kv, compressed_mask @@ -736,7 +773,12 @@ def __init__( # DeepSeek-V4 uses a mathematical attention sink (a learnable scalar per-head added to the # attention logits prior to softmax, rather than a physical key/value token). We unconditionally # initialize it here, overriding the base Attention class which disables it by default. - self.sinks = nnx.data(nnx.Param(jnp.zeros((self.num_query_heads,), dtype=self.weight_dtype), sharding=(None,))) + self.sinks = nnx.data( + nnx.Param( + jnp.zeros((self.num_query_heads,), dtype=self.weight_dtype), + sharding=(None,), + ) + ) def _init_projections(self, inputs_q_shape: Tuple, inputs_kv_shape: Tuple) -> None: """Initializes the compressed projections and Unweighted RMSNorms.""" @@ -980,7 +1022,8 @@ def __call__( 6. Flatten & Dense (o_b_proj): -> `[batch, q_length, emb_dim]`. """ q, q_normed = self.compressed_query_projection(inputs_q, inputs_positions, model_mode) - k, v = self.compressed_kv_projection(inputs_kv, inputs_positions, model_mode) + q = checkpoint_name(q, "query_proj") + kv, _ = self.compressed_kv_projection(inputs_kv, inputs_positions, model_mode) # Generate compressed representations based on the configured layer type compressed_kv = None @@ -1011,8 +1054,9 @@ def __call__( # Extend local KV tensors with the compressed blocks if compressed_kv is not None: - k = jnp.concatenate([k, compressed_kv], axis=1) - v = jnp.concatenate([v, compressed_kv], axis=1) + kv = jnp.concatenate([kv, compressed_kv], axis=1) + + kv = checkpoint_name(kv, "kv_proj") # Prepare the mask shape for the underlying AttentionOp if compressed_mask is not None: @@ -1026,8 +1070,8 @@ def __call__( # -> [batch, q_length, num_query_heads, head_dim] attn_out = self.attention_op( q, - k, - v, + kv, + kv, decoder_segment_ids, inputs_positions, model_mode, @@ -1038,6 +1082,8 @@ def __call__( # Reverse RoPE on Values attn_out = self.rotary_embedding(attn_out, inputs_positions, unsqueeze_dim=-2, reverse=True) + attn_out = checkpoint_name(attn_out, "attention_out") + # Project outputs through Grouped Linear layers b, s, h, d = attn_out.shape # -> [batch, q_length, o_groups, in_features_per_group] @@ -1051,6 +1097,7 @@ def __call__( # -> [batch, q_length, emb_dim] final_out = self.o_b_proj(grouped_flat) + final_out = checkpoint_name(final_out, "out_proj") return final_out, None diff --git a/src/maxtext/layers/decoders.py b/src/maxtext/layers/decoders.py index fc570e69bb..a915f7f85d 100644 --- a/src/maxtext/layers/decoders.py +++ b/src/maxtext/layers/decoders.py @@ -328,6 +328,7 @@ def minimal_policy(self, with_context=False, with_quantization=False): "query_proj", "value_proj", "key_proj", + "kv_proj", "qkv_proj", "out_proj", "mlpwi_0", @@ -378,6 +379,7 @@ def get_remat_policy(self): "query_proj", "value_proj", "key_proj", + "kv_proj", "qkv_proj", "context", "out_proj", @@ -387,6 +389,7 @@ def get_remat_policy(self): "query_proj", "value_proj", "key_proj", + "kv_proj", "qkv_proj", "out_proj", "mlpwo", @@ -396,6 +399,7 @@ def get_remat_policy(self): "query_proj", "value_proj", "key_proj", + "kv_proj", "qkv_proj", "out_proj", ) @@ -404,12 +408,13 @@ def get_remat_policy(self): "query_proj", "value_proj", "key_proj", + "kv_proj", "qkv_proj", ) elif cfg.remat_policy == "qkv_proj_offloaded": policy = jax.checkpoint_policies.save_and_offload_only_these_names( names_which_can_be_saved=[], - names_which_can_be_offloaded=["query_proj", "value_proj", "key_proj"], + names_which_can_be_offloaded=["query_proj", "value_proj", "key_proj", "kv_proj"], offload_src="device", offload_dst="pinned_host", ) @@ -421,6 +426,7 @@ def get_remat_policy(self): "query_proj", "value_proj", "key_proj", + "kv_proj", "qkv_proj", "out_proj", "mlpwi_0", diff --git a/src/maxtext/layers/nnx_decoders.py b/src/maxtext/layers/nnx_decoders.py index 1b14803d94..df2672be95 100644 --- a/src/maxtext/layers/nnx_decoders.py +++ b/src/maxtext/layers/nnx_decoders.py @@ -1140,6 +1140,7 @@ def minimal_policy(self, with_context=False, with_quantization=False): "query_proj", "value_proj", "key_proj", + "kv_proj", "qkv_proj", "out_proj", "mlpwi_0", @@ -1187,6 +1188,7 @@ def get_remat_policy(self): "query_proj", "value_proj", "key_proj", + "kv_proj", "qkv_proj", "context", "out_proj", @@ -1196,6 +1198,7 @@ def get_remat_policy(self): "query_proj", "value_proj", "key_proj", + "kv_proj", "qkv_proj", "out_proj", "mlpwo", @@ -1205,6 +1208,7 @@ def get_remat_policy(self): "query_proj", "value_proj", "key_proj", + "kv_proj", "qkv_proj", "out_proj", ) @@ -1213,6 +1217,7 @@ def get_remat_policy(self): "query_proj", "value_proj", "key_proj", + "kv_proj", "qkv_proj", ) elif cfg.remat_policy == "qkv_proj_offloaded": @@ -1222,6 +1227,7 @@ def get_remat_policy(self): "query_proj", "value_proj", "key_proj", + "kv_proj", ], offload_src="device", offload_dst="pinned_host", @@ -1233,6 +1239,7 @@ def get_remat_policy(self): "query_proj", "value_proj", "key_proj", + "kv_proj", "qkv_proj", "out_proj", "mlpwi_0", diff --git a/src/maxtext/utils/estimator.py b/src/maxtext/utils/estimator.py index 14ccc8c154..5bc44169fc 100644 --- a/src/maxtext/utils/estimator.py +++ b/src/maxtext/utils/estimator.py @@ -53,7 +53,12 @@ class Action(IntEnum): class RematPolicy: """RematPolicy representing different remat policy combinations""" - def __init__(self, tensor_names: list[str], tensors: dict | None = None, initial_level: Action = Action.REMAT): + def __init__( + self, + tensor_names: list[str], + tensors: dict | None = None, + initial_level: Action = Action.REMAT, + ): self.tensors = {name: initial_level for name in tensor_names} if tensors is None else tensors self.tensor_order = tensor_names @@ -114,10 +119,32 @@ def generate_priority_list(config, provided_tensor_names): keys = { (True, 1): ["context", "qkv_proj", "mlpwi", "mlpwo", "out_proj"], (True, 2): ["context", "qkv_proj", "mlpwi_0", "mlpwi_1", "mlpwo", "out_proj"], - (False, 1): ["context", "query_proj", "key_proj", "value_proj", "mlpwi", "mlpwo", "out_proj"], - (False, 2): ["context", "query_proj", "key_proj", "value_proj", "mlpwi_0", "mlpwi_1", "mlpwo", "out_proj"], + (False, 1): [ + "context", + "query_proj", + "key_proj", + "value_proj", + "kv_proj", + "mlpwi", + "mlpwo", + "out_proj", + ], + (False, 2): [ + "context", + "query_proj", + "key_proj", + "value_proj", + "kv_proj", + "mlpwi_0", + "mlpwi_1", + "mlpwo", + "out_proj", + ], } - sort_tensor_names = sorted(keys[config.fused_mlp, len(config.mlp_activations)], key=lambda x: tensor_score(x, config)) + sort_tensor_names = sorted( + keys[config.fused_mlp, len(config.mlp_activations)], + key=lambda x: tensor_score(x, config), + ) return [key for key in sort_tensor_names if key not in provided_tensor_names] @@ -150,6 +177,7 @@ def tensor_score(tensor_name: str, config) -> tuple: -config.num_query_heads * config.head_dim, ), "key_proj": (-config.emb_dim, -config.num_kv_heads * config.head_dim), + "kv_proj": (-config.emb_dim, -config.num_kv_heads * config.head_dim), "value_proj": ( -config.emb_dim, -config.num_kv_heads * config.head_dim, @@ -235,7 +263,11 @@ def largest_batch_size(base_argv, policy, min_pdb=None, max_pdb=32.0, pdb_scalar print(f"No OOM at maximum batch size {max_pdb}.") return max_pdb - low, high, result = int(min_pdb * pdb_scalar), int(max_pdb * pdb_scalar), int(min_pdb * pdb_scalar) + low, high, result = ( + int(min_pdb * pdb_scalar), + int(max_pdb * pdb_scalar), + int(min_pdb * pdb_scalar), + ) while low <= high: mid = (low + high) // 2 if mid < min_pdb: @@ -338,6 +370,7 @@ def search_policy_only( def search( tensor_names, base_argv, + *, min_pdb: float | None = None, max_pdb: float = 64.0, init_policy: RematPolicy = None, @@ -475,6 +508,7 @@ def find_remat_policy_tensor_names(base_argv): "context", "query_proj", "key_proj", + "kv_proj", "value_proj", "mlpwi_0", "mlpwi_1", @@ -533,7 +567,12 @@ def main(argv_list: Sequence[str]) -> None: # MODE 2: No batch size. Search for both batch size and policy. print("No batch size provided. Searching for max batch size and policies...") # First, find the absolute max batch size that fits *even with full remat* - max_pdb = largest_batch_size(base_argv, full_remat_policy, min_pdb=1.0 / pdb_scalar, pdb_scalar=pdb_scalar) + max_pdb = largest_batch_size( + base_argv, + full_remat_policy, + min_pdb=1.0 / pdb_scalar, + pdb_scalar=pdb_scalar, + ) # Now, search for combinations, starting from full-remat up to min_pdb suggested_list.extend( diff --git a/tests/integration/estimator_test.py b/tests/integration/estimator_test.py index 7809c25c13..53c2d3f8b6 100644 --- a/tests/integration/estimator_test.py +++ b/tests/integration/estimator_test.py @@ -83,7 +83,17 @@ def _make_base_argv(self, extra_args=None): def test_is_oom_returns_bool(self): """Verify is_oom returns a boolean for a small model with full remat.""" base_argv = self._make_base_argv() - tensor_names = ["context", "query_proj", "key_proj", "value_proj", "mlpwi_0", "mlpwi_1", "mlpwo", "out_proj"] + tensor_names = [ + "context", + "query_proj", + "key_proj", + "value_proj", + "kv_proj", + "mlpwi_0", + "mlpwi_1", + "mlpwo", + "out_proj", + ] policy = RematPolicy(tensor_names=tensor_names, initial_level=Action.REMAT) jax.clear_caches() @@ -98,7 +108,17 @@ def test_is_oom_returns_bool(self): def test_search_policy_only_small_model(self): """E2E: search_policy_only returns a valid policy for a small model.""" base_argv = self._make_base_argv() - tensor_names = ["context", "query_proj", "key_proj", "value_proj", "mlpwi_0", "mlpwi_1", "mlpwo", "out_proj"] + tensor_names = [ + "context", + "query_proj", + "key_proj", + "value_proj", + "kv_proj", + "mlpwi_0", + "mlpwi_1", + "mlpwo", + "out_proj", + ] jax.clear_caches() result = search_policy_only(tensor_names, base_argv, pdb=2.0) diff --git a/tests/unit/attention_compressed_test.py b/tests/unit/attention_compressed_test.py new file mode 100644 index 0000000000..47dda737dc --- /dev/null +++ b/tests/unit/attention_compressed_test.py @@ -0,0 +1,114 @@ +# 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. + +"""Tests for compressed attention.""" + +import unittest +import jax +import jax.numpy as jnp +from jax.sharding import Mesh +from flax import nnx +from maxtext.configs.pyconfig import initialize +from maxtext.layers.attention_compressed import CompressedAttention +from tests.utils.test_helpers import get_test_config_path + + +class CompressedAttentionTest(unittest.TestCase): + """Tests for Compressed Attention.""" + + def setUp(self): + self.config = initialize( + [ + None, + get_test_config_path(), + "model_name=deepseek4-284b", + "attention=dot_product", + "qk_rope_head_dim=16", + "v_head_dim=16", + "qk_nope_head_dim=16", + ] + ) + self.mesh = Mesh(jax.devices(), ("data",)) + + def test_compressed_attention_jaxpr_tag_counts(self): + layer = CompressedAttention( + config=self.config, + num_query_heads=4, + num_kv_heads=1, + head_dim=512, + max_target_length=128, + mesh=self.mesh, + attention_kernel="dot_product", + inputs_q_shape=(1, 32, 4096), + inputs_kv_shape=(1, 32, 4096), + compress_ratio=4, + q_lora_rank=1024, + rngs=nnx.Rngs(0), + ) + + q = jnp.ones((1, 32, 4096)) + kv = jnp.ones((1, 32, 4096)) + pos = jnp.arange(32)[None, :] + seg = jnp.zeros((1, 32), dtype=jnp.int32) + + graphdef, state = nnx.split(layer) + + def forward(state, q, kv, seg, pos): + layer = nnx.merge(graphdef, state) + return layer(q, kv, seg, pos, deterministic=True) + + jaxpr = jax.make_jaxpr(forward)(state, q, kv, seg, pos) + jaxpr_str = str(jaxpr) + + self.assertEqual(jaxpr_str.count("name=query_proj"), 1) + self.assertEqual(jaxpr_str.count("name=kv_proj"), 1) + self.assertEqual(jaxpr_str.count("name=attention_out"), 1) + self.assertEqual(jaxpr_str.count("name=out_proj"), 1) + + def test_compressed_attention_no_double_tagging(self): + layer = CompressedAttention( + config=self.config, + num_query_heads=4, + num_kv_heads=1, + head_dim=512, + max_target_length=128, + mesh=self.mesh, + attention_kernel="dot_product", + inputs_q_shape=(1, 32, 4096), + inputs_kv_shape=(1, 32, 4096), + compress_ratio=4, + q_lora_rank=1024, + rngs=nnx.Rngs(0), + ) + + q = jnp.ones((1, 32, 4096)) + kv = jnp.ones((1, 32, 4096)) + pos = jnp.arange(32)[None, :] + seg = jnp.zeros((1, 32), dtype=jnp.int32) + + graphdef, state = nnx.split(layer) + + def forward(state, q, kv, seg, pos): + layer = nnx.merge(graphdef, state) + return layer(q, kv, seg, pos, deterministic=True) + + jaxpr = jax.make_jaxpr(forward)(state, q, kv, seg, pos) + jaxpr_str = str(jaxpr) + + self.assertEqual(jaxpr_str.count("name=key_proj"), 0) + self.assertEqual(jaxpr_str.count("name=value_proj"), 0) + + +if __name__ == "__main__": + unittest.main()