Repository navigation
Conversation
There was a problem hiding this comment.
Code Review
This pull request introduces support for a per-layer context parallel (CP) strategy and a new 'halo' sliding-window CP attention mechanism in MaxText. It adds configuration options, validation logic, and the core implementation of the halo strategy using ppermute communication. Comprehensive unit tests are also added. The review comments identify two important improvements: guarding against a potential TypeError when validating sliding_window_size if it is None, and ensuring XLA_FLAGS is set before importing JAX in the test file to guarantee proper device initialization.
| if ( | ||
| context_parallel_strategy == "halo" or local_context_parallel_strategy == "halo" | ||
| ) and self.sliding_window_size <= 0: | ||
| raise ValueError("Halo context parallelism requires sliding_window_size > 0.") |
There was a problem hiding this comment.
If self.sliding_window_size is None, the comparison self.sliding_window_size <= 0 will raise a TypeError. Guarding against None values ensures robust configuration validation.
| if ( | |
| context_parallel_strategy == "halo" or local_context_parallel_strategy == "halo" | |
| ) and self.sliding_window_size <= 0: | |
| raise ValueError("Halo context parallelism requires sliding_window_size > 0.") | |
| if ( | |
| context_parallel_strategy == "halo" or local_context_parallel_strategy == "halo" | |
| ) and (self.sliding_window_size is None or self.sliding_window_size <= 0): | |
| raise ValueError("Halo context parallelism requires sliding_window_size > 0.") |
| import os | ||
| import types | ||
| from absl.testing import absltest | ||
| from absl.testing import parameterized | ||
| from flax import linen as nn | ||
| from flax import nnx | ||
| import jax | ||
| import jax.numpy as jnp | ||
| from jax.sharding import Mesh | ||
| from jax.sharding import NamedSharding | ||
| from jax.sharding import PartitionSpec as P | ||
| from maxtext.common.common_types import AttentionType | ||
| from maxtext.common.common_types import MODEL_MODE_TRAIN | ||
| from maxtext.configs import pyconfig | ||
| from maxtext.layers import attention_op as attention_op_lib | ||
| from maxtext.utils import max_utils | ||
| import numpy as np | ||
|
|
||
| os.environ.setdefault("XLA_FLAGS", "--xla_force_host_platform_device_count=8") |
There was a problem hiding this comment.
Setting XLA_FLAGS via os.environ after importing jax can be unreliable because JAX may initialize its backend and devices upon import. To guarantee that the CPU device count is correctly set to 8, os.environ should be configured before import jax.
| import os | |
| import types | |
| from absl.testing import absltest | |
| from absl.testing import parameterized | |
| from flax import linen as nn | |
| from flax import nnx | |
| import jax | |
| import jax.numpy as jnp | |
| from jax.sharding import Mesh | |
| from jax.sharding import NamedSharding | |
| from jax.sharding import PartitionSpec as P | |
| from maxtext.common.common_types import AttentionType | |
| from maxtext.common.common_types import MODEL_MODE_TRAIN | |
| from maxtext.configs import pyconfig | |
| from maxtext.layers import attention_op as attention_op_lib | |
| from maxtext.utils import max_utils | |
| import numpy as np | |
| os.environ.setdefault("XLA_FLAGS", "--xla_force_host_platform_device_count=8") | |
| import os | |
| os.environ.setdefault("XLA_FLAGS", "--xla_force_host_platform_device_count=8") | |
| import types | |
| from absl.testing import absltest | |
| from absl.testing import parameterized | |
| from flax import linen as nn | |
| from flax import nnx | |
| import jax | |
| import jax.numpy as jnp | |
| from jax.sharding import Mesh | |
| from jax.sharding import NamedSharding | |
| from jax.sharding import PartitionSpec as P | |
| from maxtext.common.common_types import AttentionType | |
| from maxtext.common.common_types import MODEL_MODE_TRAIN | |
| from maxtext.configs import pyconfig | |
| from maxtext.layers import attention_op as attention_op_lib | |
| from maxtext.utils import max_utils | |
| import numpy as np |
| ### Determine if we want to use load balance for context parallelism | ||
| context_parallel_load_balance: true | ||
| context_parallel_strategy: "all_gather" # "all_gather", "ring", "ulysses", or "usp" | ||
| context_parallel_strategy: "all_gather" # "all_gather", "ring", "ulysses", "usp", or "halo" |
There was a problem hiding this comment.
We can remove halo here
| import numpy as np | ||
| from tokamax._src.ops.attention import base as tokamax_attention_base | ||
| from tokamax._src.ops.attention import pallas_triton as tokamax_pallas_triton | ||
| from tokamax._src.ops.experimental.tpu.splash_attention import splash_attention_kernel as tokamax_splash_kernel | ||
| from tokamax._src.ops.experimental.tpu.splash_attention import splash_attention_mask as tokamax_splash_mask | ||
| try: | ||
| from tokamax._src.ops.attention import base as tokamax_attention_base | ||
| from tokamax._src.ops.attention import pallas_triton as tokamax_pallas_triton | ||
| from tokamax._src.ops.experimental.tpu.splash_attention import splash_attention_kernel as tokamax_splash_kernel |
| raise ValueError("context_parallel_strategy must be one of 'all_gather', 'ring', 'ulysses', or 'usp'.") | ||
| if context_parallel_strategy not in ("all_gather", "ring", "ulysses", "usp", "halo"): | ||
| raise ValueError("context_parallel_strategy must be one of 'all_gather', 'ring', 'ulysses', 'usp', or 'halo'.") | ||
| self.context_parallel_strategy = context_parallel_strategy |
There was a problem hiding this comment.
make it simple, like I mentioned, keep the local only to halo..
| single_head_mask = mask # tokamax now just uses a single mask and assumes broadcast to all heads | ||
| if self.config.use_max_logit_estimate > 0: | ||
| sa_config = dataclasses.replace(sa_config, max_logit_const=self.config.use_max_logit_estimate) | ||
| if use_halo and jax.default_backend() == "cpu": |
| return splash_kernel | ||
|
|
||
| segment_axis_names_splash_kernel = self._logical_to_mesh_axes((Q_LENGTH,)) | ||
| segment_axis_names_splash_kernel = ( |
There was a problem hiding this comment.
better seperate the halo attention, it will cause less changes and will be more clear and readable
| ) | ||
| use_dq_carry = config.dq_reduction_steps == 3 | ||
| dq_accum = jnp.zeros((3, *q.shape), dtype=jnp.float32) if use_dq_carry else jnp.zeros(q.shape, dtype=jnp.float32) | ||
| dq_accum = jnp.zeros((3, *q.shape), dtype=q.dtype) if use_dq_carry else jnp.zeros(q.shape, dtype=jnp.float32) |
There was a problem hiding this comment.
should be fp32 always, revert
| mask_value=mask_value, | ||
| is_mqa=is_mqa, | ||
| config=config, | ||
| config=dataclasses.replace(config, residual_checkpoint_name=None), |
There was a problem hiding this comment.
why is this needed, if needed, do a proper fix.
d34004f to
89f93af
Compare
…ange for LOCAL_SLIDING
Adds local_context_parallel_strategy ('', 'all_gather', 'halo') so hybrid models like Gemma4 can override context_parallel_strategy on LOCAL_SLIDING layers (e.g., global layers using ring/ulysses/all_gather while local layers use halo). Separates the halo KV exchange path into AttentionOp.tpu_halo_flash_attention.
89f93af to
825cbdd
Compare
Summary
Adds
local_context_parallel_strategy("","all_gather","halo") so hybrid models with bothGLOBALandLOCAL_SLIDINGattention layers (such as Gemma4-26B) can overridecontext_parallel_strategyonLOCAL_SLIDINGlayers, and implements appermute-based halo KV exchange inAttentionOp.tpu_halo_flash_attentionforLOCAL_SLIDINGcontext parallelism.local_context_parallel_strategy):AttentionOp.__init__resolvesself.context_parallel_strategyusinglocal_context_parallel_strategywhenself.attention_type == AttentionType.LOCAL_SLIDING, allowingGLOBALlayers to useulysses,ring, orall_gatherwhileLOCAL_SLIDINGlayers usehaloorall_gather.AttentionOp.tpu_halo_flash_attention): Instead of all-gathering and unpermuting the full[B, S, H_kv, D]K/V sequence across thecontextaxis onLOCAL_SLIDINGlayers, each CP rank exchanges only the block-aligned sliding-window tail (halo_pad_width) of the preceding chunk vialax.ppermute. Supports both contiguous (context_parallel_load_balance=False) andDUAL_CHUNK_SWAP(context_parallel_load_balance=True) layouts without splitting the batch dimension.Related PRs
gmm_v2backwarddlhskernel)