From 5f1086085b7afce40b340e4f678e94db1b5a3b33 Mon Sep 17 00:00:00 2001 From: andrewwhitecdw Date: Wed, 12 Aug 2026 15:03:41 -0500 Subject: [PATCH] fix: make_mask uses deprecated jnp.bool instead of jnp.bool_ - Replace jnp.bool with jnp.bool_ in the bottom-right sliding-window mask path of make_mask. - Strengthen the regression test to capture the dtype actually passed to make_swa_mask, so reverting to jnp.bool fails the assertion. - No other jnp.bool call sites remain in tests/jax. Signed-off-by: Andrew White --- tests/jax/test_fused_attn.py | 33 ++++++++++++++++++++++++++++++++- 1 file changed, 32 insertions(+), 1 deletion(-) diff --git a/tests/jax/test_fused_attn.py b/tests/jax/test_fused_attn.py index 1dd92fe181..40586e376b 100644 --- a/tests/jax/test_fused_attn.py +++ b/tests/jax/test_fused_attn.py @@ -3,6 +3,7 @@ # See LICENSE for license information. """Tests for fused attention""" import os +import sys from enum import Enum, auto from dataclasses import dataclass, field from functools import partial @@ -224,7 +225,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 +237,36 @@ def make_mask( return mask +def test_make_mask_bottom_right_swa_dtype(monkeypatch): + """Regression test: make_mask bottom-right branch should use jnp.bool_, not jnp.bool.""" + captured_dtypes = [] + mod = sys.modules[make_mask.__module__] + original_make_swa_mask = mod.make_swa_mask + + def _captured_make_swa_mask(*args, **kwargs): + captured_dtypes.append(kwargs.get("dtype")) + return original_make_swa_mask(*args, **kwargs) + + monkeypatch.setattr(mod, "make_swa_mask", _captured_make_swa_mask) + + 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_ + assert ( + jnp.bool_ in captured_dtypes + ), "bottom-right sliding-window mask should be built with jnp.bool_" + + @jax.jit def get_seqlens_and_offsets(segment_ids): batch, max_seqlen = segment_ids.shape