diff --git a/tests/jax/test_fused_attn.py b/tests/jax/test_fused_attn.py index 1dd92fe181..b7feabcf1d 100644 --- a/tests/jax/test_fused_attn.py +++ b/tests/jax/test_fused_attn.py @@ -224,7 +224,7 @@ def make_mask( segment_pos_q, segment_pos_kv, window_size, - dtype=jnp.bool, + dtype=jnp.bool_, segment_ids_q=segment_ids_q, segment_ids_kv=segment_ids_kv, ) @@ -236,6 +236,23 @@ def make_mask( return mask +def test_make_mask_bottom_right_swa_dtype(): + """Regression test: make_mask bottom-right branch should use jnp.bool_, not jnp.bool.""" + batch, seqlen = 2, 16 + segment_ids = jnp.ones((batch, seqlen), dtype=jnp.int32) + segment_pos = jnp.broadcast_to(jnp.arange(seqlen, dtype=jnp.int32), (batch, seqlen)) + window_size = (4, 0) + mask = make_mask( + segment_ids, + segment_ids, + segment_pos, + segment_pos, + AttnMaskType.PADDING_CAUSAL_BOTTOM_RIGHT_MASK, + window_size, + ) + assert mask.dtype == jnp.bool_ + + @jax.jit def get_seqlens_and_offsets(segment_ids): batch, max_seqlen = segment_ids.shape