@@ -726,109 +726,6 @@ def make_ring_attention(
726726# ---------------------------------------------------------------------------
727727
728728
729- def _custom_bidirectional_ring_forward (
730- q : jax .Array ,
731- k : jax .Array ,
732- v : jax .Array ,
733- * ,
734- block_sizes : "custom_splash._BlockSizes" ,
735- bkv_compute_in : int ,
736- orig_q_seq_len : int ,
737- orig_kv_seq_len : int ,
738- use_base2_exp : bool ,
739- use_experimental_scheduler : bool ,
740- vmem_limit_bytes : int | None ,
741- mask_value : float ,
742- ring_axis : str ,
743- ) -> jax .Array :
744- """Wrap-free (bidirectional) ring attention for a NON-wrapping ring axis.
745-
746- On a torus dimension the +1-mod-R ppermute is nearest-neighbor, but on a
747- cut/non-wrapping axis (e.g. the size-4 z line of a v7x-16 slice) the wrap edge
748- (R-1 -> 0) spans the whole line diameter. Instead of one rotating stream with
749- that long, congested wrap, stream K/V BOTH directions one hop at a time:
750- - rightward stream: device i holds KV_{i-t} after t hops,
751- - leftward stream: device i holds KV_{i+t} after t hops,
752- with out-of-range shards (line ends) masked out of the online softmax. Every
753- step is a single hop and uses both link directions; no edge ever spans the
754- diameter.
755-
756- Trade-off: each device computes ~2x attention blocks (the line-end ones are
757- masked), traded for the removed multi-hop wrap. Net win when the ring is
758- comms-bound (the case on a non-wrapping axis). Operates on the FULL real ring
759- axis (no sub-group perm).
760- """
761- axis_size = lax .axis_size (ring_axis )
762- idx = lax .axis_index (ring_axis )
763- exp_fn = jnp .exp2 if use_base2_exp else jnp .exp
764-
765- def _attn (kc , vc ):
766- o , m , l = custom_splash ._splash_attention_forward_ring ( # pylint: disable=protected-access
767- q ,
768- kc ,
769- vc ,
770- block_sizes ,
771- bkv_compute_in ,
772- q_seq_len = orig_q_seq_len ,
773- kv_seq_len = orig_kv_seq_len ,
774- use_base2_exp = use_base2_exp ,
775- use_experimental_scheduler = use_experimental_scheduler ,
776- vmem_limit_bytes = vmem_limit_bytes ,
777- )
778- return o .astype (jnp .float32 ), m .astype (jnp .float32 ), l .astype (jnp .float32 )
779-
780- def _merge (m , l , o , mc , lc , oc , valid ):
781- # Nullify invalid (line-end) contributions. Force mc to mask_value (beta -> 0)
782- # AND zero lc/oc so a non-finite zero-buffer result can't leak via 0*inf=nan.
783- mc = jnp .where (valid , mc , mask_value )
784- lc = jnp .where (valid , lc , 0.0 )
785- oc = jnp .where (valid , oc , 0.0 )
786- m_next = jnp .maximum (m , mc )
787- alpha = exp_fn (m - m_next )
788- beta = exp_fn (mc - m_next )
789- return m_next , alpha * l + beta * lc , alpha [..., None ] * o + beta [..., None ] * oc
790-
791- # t=0: own shard (always valid). _attn returns (o, m, l).
792- o , m , l = _attn (k , v )
793-
794- # Non-wrapping one-hop shifts (line ends send/receive nothing).
795- shift_r = partial (lax .ppermute , axis_name = ring_axis , perm = [(i , i + 1 ) for i in range (axis_size - 1 )])
796- shift_l = partial (lax .ppermute , axis_name = ring_axis , perm = [(i , i - 1 ) for i in range (1 , axis_size )])
797-
798- # Prime buffers for t=1 (one hop each direction): device i -> KV_{i-1}, KV_{i+1}.
799- kr , vr = shift_r (k ), shift_r (v )
800- kl , vl = shift_l (k ), shift_l (v )
801-
802- def body (carry , t ):
803- m , l , o , kr , vr , kl , vl = carry
804- valid_r = (idx - t ) >= 0
805- valid_l = (idx + t ) <= (axis_size - 1 )
806- # Feed real (own) K/V on invalid steps so _attn never runs on a degenerate
807- # zero buffer (line ends receive 0 from the partial ppermute); masked below.
808- kr_s , vr_s = jnp .where (valid_r , kr , k ), jnp .where (valid_r , vr , v )
809- kl_s , vl_s = jnp .where (valid_l , kl , k ), jnp .where (valid_l , vl , v )
810- # Compute against the current shards (KV_{i-t}, KV_{i+t}) ...
811- o_r , m_r , l_r = _attn (kr_s , vr_s )
812- m , l , o = _merge (m , l , o , m_r , l_r , o_r , valid_r )
813- o_l , m_l , l_l = _attn (kl_s , vl_s )
814- m , l , o = _merge (m , l , o , m_l , l_l , o_l , valid_l )
815- # ... and prefetch the next hop (independent of the matmuls above -> overlaps).
816- kr_n , vr_n = shift_r (kr ), shift_r (vr )
817- kl_n , vl_n = shift_l (kl ), shift_l (vl )
818- return (m , l , o , kr_n , vr_n , kl_n , vl_n ), None
819-
820- (_ , l_final , o_final , * _ ), _ = lax .scan (
821- body ,
822- (m , l , o , kr , vr , kl , vl ),
823- xs = jnp .arange (1 , axis_size ),
824- length = axis_size - 1 ,
825- unroll = True ,
826- )
827-
828- l_inv = jnp .where (l_final == 0.0 , 0.0 , 1.0 / l_final )
829- return (o_final * l_inv [..., None ]).astype (q .dtype )
830-
831-
832729def _custom_ring_attention_forward (
833730 q : jax .Array ,
834731 k : jax .Array ,
@@ -845,7 +742,6 @@ def _custom_ring_attention_forward(
845742 ring_axis : str ,
846743 ring_size : int | None = None ,
847744 perm : list [tuple [int , int ]] | None = None ,
848- bidirectional : bool = False ,
849745) -> jax .Array :
850746 """Forward-only ring attention using the custom dense splash kernel.
851747
@@ -880,26 +776,6 @@ def _custom_ring_attention_forward(
880776 Normalized attention output, shape `(num_q_heads, q_seq_len, head_dim_v)`.
881777 """
882778 axis_size = lax .axis_size (ring_axis )
883- if bidirectional :
884- if perm is not None or (ring_size is not None and ring_size != axis_size ):
885- raise ValueError (
886- "bidirectional (wrap-free) ring requires perm=None and ring_size==axis_size "
887- "(it operates on the full real ring axis)."
888- )
889- return _custom_bidirectional_ring_forward (
890- q ,
891- k ,
892- v ,
893- block_sizes = block_sizes ,
894- bkv_compute_in = bkv_compute_in ,
895- orig_q_seq_len = orig_q_seq_len ,
896- orig_kv_seq_len = orig_kv_seq_len ,
897- use_base2_exp = use_base2_exp ,
898- use_experimental_scheduler = use_experimental_scheduler ,
899- vmem_limit_bytes = vmem_limit_bytes ,
900- mask_value = mask_value ,
901- ring_axis = ring_axis ,
902- )
903779 if ring_size is None :
904780 ring_size = axis_size
905781 if perm is None :
@@ -971,7 +847,6 @@ def make_custom_ring_attention(
971847 ring_axis : str = "context" ,
972848 ring_size : int | None = None ,
973849 perm : list [tuple [int , int ]] | None = None ,
974- bidirectional : bool = False ,
975850):
976851 """Builds a forward-only ring-attention callable around the custom kernel.
977852
@@ -983,10 +858,6 @@ def make_custom_ring_attention(
983858 `ring_size` / `perm` let a caller restrict the rotation to a ring sub-group of
984859 the axis (for the hybrid Ulysses+Ring / USP split); when omitted the rotation
985860 covers the whole `ring_axis`.
986-
987- `bidirectional=True` selects the wrap-free schedule (streams K/V both directions
988- one hop at a time) for a NON-wrapping ring axis, avoiding the diameter-length
989- wrap hop. Requires `perm=None` and the full real ring axis (no sub-group).
990861 """
991862
992863 def _ring (q , k , v ):
@@ -1005,7 +876,6 @@ def _ring(q, k, v):
1005876 ring_axis = ring_axis ,
1006877 ring_size = ring_size ,
1007878 perm = perm ,
1008- bidirectional = bidirectional ,
1009879 )
1010880
1011881 return _ring
0 commit comments