Skip to content

Attention: per-layer local_context_parallel_strategy and halo KV exchange for LOCAL_SLIDING - #5608

Draft
csgoogle wants to merge 1 commit into
AI-Hypercomputer:mainfrom
csgoogle:g4-v6e128-attention
Draft

csgoogle wants to merge 1 commit into
AI-Hypercomputer:mainfrom
csgoogle:g4-v6e128-attention

Conversation

@csgoogle

@csgoogle csgoogle commented Oct 8, 2026 •

Copy link
Copy Markdown
Contributor

Summary

Adds local_context_parallel_strategy ("", "all_gather", "halo") so hybrid models with both GLOBAL and LOCAL_SLIDING attention layers (such as Gemma4-26B) can override context_parallel_strategy on LOCAL_SLIDING layers, and implements a ppermute-based halo KV exchange in AttentionOp.tpu_halo_flash_attention for LOCAL_SLIDING context parallelism.

  • Per-layer CP strategy (local_context_parallel_strategy): AttentionOp.__init__ resolves self.context_parallel_strategy using local_context_parallel_strategy when self.attention_type == AttentionType.LOCAL_SLIDING, allowing GLOBAL layers to use ulysses, ring, or all_gather while LOCAL_SLIDING layers use halo or all_gather.
  • Separated Halo Flash Attention (AttentionOp.tpu_halo_flash_attention): Instead of all-gathering and unpermuting the full [B, S, H_kv, D] K/V sequence across the context axis on LOCAL_SLIDING layers, each CP rank exchanges only the block-aligned sliding-window tail (halo_pad_width) of the preceding chunk via lax.ppermute. Supports both contiguous (context_parallel_load_balance=False) and DUAL_CHUNK_SWAP (context_parallel_load_balance=True) layouts without splitting the batch dimension.

Related PRs

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment thread src/maxtext/configs/types.py Outdated
Comment on lines +4874 to +4877
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.")

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

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.

Suggested change
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.")

Comment on lines +19 to +37
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")

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

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.

Suggested change
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

Comment thread src/maxtext/configs/base.yml Outdated
### 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"

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We can remove halo here

Comment thread src/maxtext/layers/attention_op.py Outdated
Comment on lines +77 to +81
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

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

revert

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

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

make it simple, like I mentioned, keep the local only to halo..

Comment thread src/maxtext/layers/attention_op.py Outdated
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":

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

remove

Comment thread src/maxtext/layers/attention_op.py Outdated
return splash_kernel

segment_axis_names_splash_kernel = self._logical_to_mesh_axes((Q_LENGTH,))
segment_axis_names_splash_kernel = (

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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)

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

should be fp32 always, revert

mask_value=mask_value,
is_mqa=is_mqa,
config=config,
config=dataclasses.replace(config, residual_checkpoint_name=None),

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

why is this needed, if needed, do a proper fix.

@csgoogle
csgoogle force-pushed the g4-v6e128-attention branch 3 times, most recently from d34004f to 89f93af Compare October 8, 2026 17:46
@csgoogle csgoogle changed the title Attention: per-layer Halo context parallelism for LOCAL_SLIDING layers & Tokamax Ring Attention fixes Attention: per-layer local_context_parallel_strategy and halo KV exchange for LOCAL_SLIDING Oct 8, 2026
…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.
@csgoogle
csgoogle force-pushed the g4-v6e128-attention branch from 89f93af to 825cbdd Compare October 8, 2026 18:03

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant