Download src/hamiltonzero/evaluation/sequence_parallel.py from Twobombs/HamiltonZero: direct link, hf CLI and curl.
- Browser
- Download file 5.8 kB
-
https://huggingface.co/Twobombs/HamiltonZero/resolve/main/src/hamiltonzero/evaluation/sequence_parallel.py
- Command line
-
hf download hf://Twobombs/HamiltonZero/src/hamiltonzero/evaluation/sequence_parallel.py
-
curl -L -o sequence_parallel.py https://huggingface.co/Twobombs/HamiltonZero/resolve/main/src/hamiltonzero/evaluation/sequence_parallel.py
5.8 kB
| # Copyright (c) 2026 Simulacra Research Inc. | |
| # SPDX-License-Identifier: Apache-2.0 | |
| from __future__ import annotations | |
| import jax | |
| import jax.numpy as jnp | |
| from .pallas_mha import BlockSizes, noncausal_bias_mha | |
| def _validate_rectangular_attention_inputs( | |
| q_local: jax.Array, | |
| k_global: jax.Array, | |
| v_global: jax.Array, | |
| edge_bias_local: jax.Array, | |
| key_mask: jax.Array, | |
| ) -> None: | |
| if q_local.ndim != 3 or k_global.ndim != 3 or v_global.ndim != 3: | |
| raise ValueError("q, k, and v must have shapes [sequence, heads, head_dim]") | |
| q_len, n_heads, head_dim = q_local.shape | |
| kv_len = k_global.shape[0] | |
| if q_len < 1 or kv_len < 1: | |
| raise ValueError("query and key sequence lengths must both be positive") | |
| if k_global.shape != v_global.shape: | |
| raise ValueError( | |
| f"k and v shapes must match; got {k_global.shape} and {v_global.shape}" | |
| ) | |
| if k_global.shape[1:] != (n_heads, head_dim): | |
| raise ValueError( | |
| "q, k, and v must have the same head count and head dimension; " | |
| f"got {q_local.shape}, {k_global.shape}, and {v_global.shape}" | |
| ) | |
| if edge_bias_local.shape != (q_len, kv_len, n_heads): | |
| raise ValueError( | |
| "edge bias must have shape [local_queries, global_keys, heads]; " | |
| f"got {edge_bias_local.shape}, expected {(q_len, kv_len, n_heads)}" | |
| ) | |
| if key_mask.shape != (kv_len,): | |
| raise ValueError(f"key mask must have shape {(kv_len,)}, got {key_mask.shape}") | |
| arrays = (q_local, k_global, v_global, edge_bias_local) | |
| if any(value.dtype != jnp.float32 for value in arrays): | |
| raise TypeError( | |
| "large-N rectangular attention is fp32-only; got " | |
| + ", ".join(str(value.dtype) for value in arrays) | |
| ) | |
| def _dividing_block_size(length: int, requested: int | None) -> int: | |
| block = min(length, 128 if requested is None else requested) | |
| if block < 1: | |
| raise ValueError(f"block size must be positive, got {block}") | |
| while length % block: | |
| block //= 2 | |
| return block | |
| def pallas_rectangular_edge_attention( | |
| q_local: jax.Array, | |
| k_global: jax.Array, | |
| v_global: jax.Array, | |
| edge_bias_local: jax.Array, | |
| key_mask: jax.Array, | |
| *, | |
| block_k: int | None = None, | |
| ) -> jax.Array: | |
| _validate_rectangular_attention_inputs( | |
| q_local, k_global, v_global, edge_bias_local, key_mask | |
| ) | |
| sm_scale = float(q_local.shape[-1]) ** -0.5 | |
| q_len = q_local.shape[0] | |
| kv_len = k_global.shape[0] | |
| bq = _dividing_block_size(q_len, None) | |
| bk = _dividing_block_size(kv_len, block_k) | |
| block_sizes = BlockSizes(block_q=bq, block_k=bk) | |
| masked_bias = jnp.where( | |
| key_mask[None, :, None].astype(bool), | |
| edge_bias_local, | |
| jnp.asarray(-1.0e30, dtype=jnp.float32), | |
| ) | |
| return noncausal_bias_mha( | |
| (q_local * jnp.asarray(sm_scale, dtype=jnp.float32))[None], | |
| k_global[None], | |
| v_global[None], | |
| masked_bias[None], | |
| block_sizes=block_sizes, | |
| )[0] | |
| def ring_learned_fwl2_local( | |
| a_local: jax.Array, | |
| b_local: jax.Array, | |
| *, | |
| axis_name: str, | |
| axis_size: int, | |
| ) -> jax.Array: | |
| local_rows, global_columns, channels = a_local.shape | |
| if b_local.shape != (local_rows, global_columns, channels): | |
| raise ValueError( | |
| f"local 2-FWL shapes must match; got {a_local.shape}, {b_local.shape}" | |
| ) | |
| return ring_learned_fwl2_columns_local( | |
| a_local, | |
| b_local, | |
| axis_name=axis_name, | |
| axis_size=axis_size, | |
| ) | |
| def ring_learned_fwl2_columns_local( | |
| a_local: jax.Array, | |
| b_local_columns: jax.Array, | |
| *, | |
| axis_name: str, | |
| axis_size: int, | |
| ) -> jax.Array: | |
| if a_local.ndim != 3 or b_local_columns.ndim != 3: | |
| raise ValueError( | |
| "local 2-FWL operands must both be rank three; got " | |
| f"{a_local.shape} and {b_local_columns.shape}" | |
| ) | |
| local_rows, global_columns, channels = a_local.shape | |
| if b_local_columns.shape[0] != local_rows: | |
| raise ValueError( | |
| "local A/B row counts must match; got " | |
| f"{local_rows} and {b_local_columns.shape[0]}" | |
| ) | |
| if b_local_columns.shape[2] != channels: | |
| raise ValueError( | |
| "local A/B channel counts must match; got " | |
| f"{channels} and {b_local_columns.shape[2]}" | |
| ) | |
| if a_local.dtype != b_local_columns.dtype: | |
| raise TypeError( | |
| "local A/B dtypes must match; got " | |
| f"{a_local.dtype} and {b_local_columns.dtype}" | |
| ) | |
| if global_columns != local_rows * axis_size: | |
| raise ValueError( | |
| "2-FWL row shards must evenly tile the contracted axis; " | |
| f"got local_rows={local_rows}, columns={global_columns}, " | |
| f"axis_size={axis_size}" | |
| ) | |
| def contribution(b_panel, origin): | |
| a_panel = jax.lax.dynamic_slice_in_dim( | |
| a_local, origin * local_rows, local_rows, axis=1 | |
| ) | |
| return jnp.einsum("ikc,kjc->ijc", a_panel, b_panel) | |
| origin0 = jax.lax.axis_index(axis_name).astype(jnp.int32) | |
| accumulator0 = contribution(b_local_columns, origin0) | |
| ring_permutation = [(lane, (lane + 1) % axis_size) for lane in range(axis_size)] | |
| def ring_step(carry, _): | |
| b_panel, origin, accumulator = carry | |
| b_panel = jax.lax.ppermute(b_panel, axis_name=axis_name, perm=ring_permutation) | |
| origin = (origin - jnp.asarray(1, jnp.int32)) % axis_size | |
| accumulator = accumulator + contribution(b_panel, origin) | |
| return (b_panel, origin, accumulator), None | |
| (_, _, result), _ = jax.lax.scan( | |
| ring_step, | |
| (b_local_columns, origin0, accumulator0), | |
| xs=None, | |
| length=axis_size - 1, | |
| ) | |
| return result | |