From 61685d7b1e025f5cd46c3faecfdcb4b086bad767 Mon Sep 17 00:00:00 2001 From: Kshitij Lakhani Date: Tue, 9 Jun 2026 23:08:02 -0700 Subject: [PATCH 1/9] Add support in lower level JAX API for returning max logit and softmax aux to the user from TE JAX fused attn output Signed-off-by: Kshitij Lakhani --- transformer_engine/jax/attention.py | 74 +++++++++++++++++++++++++++-- 1 file changed, 69 insertions(+), 5 deletions(-) diff --git a/transformer_engine/jax/attention.py b/transformer_engine/jax/attention.py index ecca4a3871..9d3aaae90d 100644 --- a/transformer_engine/jax/attention.py +++ b/transformer_engine/jax/attention.py @@ -339,6 +339,7 @@ def is_fused_attn_kernel_available( head_dim_qk, head_dim_v, window_size: Optional[Tuple[int, int]] = None, + return_max_logit: bool = False, ): """ To check whether the fused attention kernel is supported @@ -362,6 +363,7 @@ def make_helper(attn_mask_type): head_dim_qk, head_dim_v, window_size_tuple, + return_max_logit, ) return make_helper(attn_mask_type).is_fused_attn_kernel_available() @@ -1053,6 +1055,8 @@ def _legacy_fused_attn( context_parallel_causal_load_balanced: bool = False, context_parallel_axis: str = "", softmax_offset: Optional[jnp.ndarray] = None, + return_max_logit: bool = False, + return_softmax_aux: bool = False, ): """ Perform non-THD (non-packed) cuDNN fused attention. @@ -1084,8 +1088,18 @@ def _legacy_fused_attn( context_parallel_causal_load_balanced (bool): Indicates the sequences are ordered for causal mask load balancing when running context parallelism. context_parallel_axis (str): The name of the context parallel axis. + softmax_offset (Optional[jnp.ndarray]): An optional learnable softmax offset tensor with shape + [1, num_heads, 1, 1]. Used when softmax_type is AttnSoftmaxType.LEARNABLE_SOFTMAX. + return_max_logit (bool): If True, also return per-head maximum attention logits + in an auxiliary dictionary under ``"max_logit"``. + return_softmax_aux (bool): If True, also return backend-specific softmax statistics + in an auxiliary dictionary under ``"softmax_aux"``. Returns: - (jnp.ndarray): The output tensor from the fused attention. + jnp.ndarray: + Attention output when neither ``return_max_logit`` nor ``return_softmax_aux`` is True. + tuple[jnp.ndarray, dict[str, jnp.ndarray]]: + ``(output, aux)`` when either flag is True. ``aux`` may contain: + ``"max_logit"`` (shape ``[h]``) and/or ``"softmax_aux"`` (float32). """ assert ( not qkv_layout.is_thd() @@ -1139,6 +1153,8 @@ def _legacy_fused_attn( context_parallel_strategy=context_parallel_strategy, context_parallel_causal_load_balanced=context_parallel_causal_load_balanced, context_parallel_axis=context_parallel_axis, + return_max_logit=return_max_logit, + return_softmax_aux=return_softmax_aux, ) return output @@ -1164,6 +1180,8 @@ def fused_attn_thd( context_parallel_causal_load_balanced: bool = False, context_parallel_axis: str = "", softmax_offset: Optional[jnp.ndarray] = None, + return_max_logit: bool = False, + return_softmax_aux: bool = False, ): """ Deprecated THD fused attn, please use fusd_attn with SequenceDescriptor @@ -1218,12 +1236,17 @@ def fused_attn_thd( context_parallel_strategy=context_parallel_strategy, context_parallel_causal_load_balanced=context_parallel_causal_load_balanced, context_parallel_axis=context_parallel_axis, + return_max_logit=return_max_logit, + return_softmax_aux=return_softmax_aux, ) return output -@partial(jax.custom_vjp, nondiff_argnums=(5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18)) +@partial( + jax.custom_vjp, + nondiff_argnums=(5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20), +) def _fused_attn( qkv: Tuple[jnp.ndarray, ...], bias: Optional[jnp.ndarray], @@ -1244,6 +1267,8 @@ def _fused_attn( context_parallel_axis: str, context_checkpoint_name: str = "context", stripe_size: int | None = None, + return_max_logit: bool = False, + return_softmax_aux: bool = False, ): output, _ = _fused_attn_fwd_rule( qkv, @@ -1265,6 +1290,8 @@ def _fused_attn( context_parallel_axis, context_checkpoint_name=context_checkpoint_name, stripe_size=stripe_size, + return_max_logit=return_max_logit, + return_softmax_aux=return_softmax_aux, ) return output @@ -1289,8 +1316,10 @@ def _fused_attn_fwd_rule( context_parallel_axis, context_checkpoint_name, stripe_size, + return_max_logit, + return_softmax_aux, ): - output, softmax_aux, rng_state = tex.fused_attn_fwd( + output, softmax_aux, rng_state, max_logit = tex.fused_attn_fwd( qkv, bias, softmax_offset, @@ -1309,11 +1338,16 @@ def _fused_attn_fwd_rule( context_parallel_causal_load_balanced=context_parallel_causal_load_balanced, context_parallel_axis=context_parallel_axis, stripe_size=stripe_size, + return_max_logit=return_max_logit, ) output = checkpoint_name(output, context_checkpoint_name) softmax_aux = checkpoint_name(softmax_aux, context_checkpoint_name) rng_state = checkpoint_name(rng_state, context_checkpoint_name) - return output, ( + max_logit = checkpoint_name(max_logit, context_checkpoint_name) + attn_output = _resolve_fused_attn_output( + output, max_logit, softmax_aux, return_max_logit, return_softmax_aux + ) + return attn_output, ( qkv, bias, sequence_descriptor, @@ -1339,10 +1373,14 @@ def _fused_attn_bwd_rule( context_parallel_axis, context_checkpoint_name, stripe_size, + return_max_logit, + return_softmax_aux, ctx, dz, ): del context_checkpoint_name + if return_max_logit or return_softmax_aux: + dz, _ = dz ( qkv, bias, @@ -1388,6 +1426,18 @@ def _fused_attn_bwd_rule( ) +def _resolve_fused_attn_output(output, max_logit, softmax_aux, return_max_logit, return_softmax_aux): + if not return_max_logit and not return_softmax_aux: + return output + + aux = {} + if return_max_logit: + aux["max_logit"] = max_logit + if return_softmax_aux: + aux["softmax_aux"] = softmax_aux + return output, aux + + _fused_attn.defvjp(_fused_attn_fwd_rule, _fused_attn_bwd_rule) @@ -1468,6 +1518,8 @@ def fused_attn( score_mod_bprop: Optional[Callable] = None, score_mod_tensors: Optional[Mapping[str, Any]] = None, score_mod_bprop_tensors: Optional[Mapping[str, Any]] = None, + return_max_logit: bool = False, + return_softmax_aux: bool = False, ): """ Perform cuDNN fused attention. @@ -1524,8 +1576,16 @@ def fused_attn( non-differentiable auxiliary inputs. score_mod_bprop_tensors (Optional[Mapping[str, Any]]): Additional tensors or Python/NumPy scalars made available to `score_mod_bprop`. + return_max_logit (bool): If True, also return per-head maximum attention logits + in an auxiliary dictionary under ``"max_logit"``. + return_softmax_aux (bool): If True, also return backend-specific softmax statistics + in an auxiliary dictionary under ``"softmax_aux"``. Returns: - (jnp.ndarray): The output tensor from the fused attention. + jnp.ndarray: + Attention output when neither ``return_max_logit`` nor ``return_softmax_aux`` is True. + tuple[jnp.ndarray, dict[str, jnp.ndarray]]: + ``(output, aux)`` when either flag is True. ``aux`` may contain: + ``"max_logit"`` (shape ``[h]``) and/or ``"softmax_aux"`` (float32). Examples (non-THD, also known as non-packed): >>> # q_segment_ids = [[1, 1, 1, 0], [1, 1, 0, 0]], 0 means padded tokens @@ -1628,6 +1688,8 @@ def fused_attn( context_parallel_causal_load_balanced=context_parallel_causal_load_balanced, context_parallel_axis=context_parallel_axis, softmax_offset=softmax_offset, + return_max_logit=return_max_logit, + return_softmax_aux=return_softmax_aux, ) if max_segments_per_seq > 1 and not qkv_layout.is_thd(): warnings.warn( @@ -1658,5 +1720,7 @@ def fused_attn( context_parallel_axis=context_parallel_axis, context_checkpoint_name=context_checkpoint_name, stripe_size=stripe_size, + return_max_logit=return_max_logit, + return_softmax_aux=return_softmax_aux, ) return output From c23bc79a8430e2a6bb88314fcbc78bb2ee85a7d4 Mon Sep 17 00:00:00 2001 From: Kshitij Lakhani Date: Tue, 9 Jun 2026 23:10:57 -0700 Subject: [PATCH 2/9] Add support for returning reduced per head max logit. Plumb max logit and softmax through the JAX fused attn primitives Signed-off-by: Kshitij Lakhani --- .../jax/cpp_extensions/attention.py | 129 ++++++++++++++---- 1 file changed, 104 insertions(+), 25 deletions(-) diff --git a/transformer_engine/jax/cpp_extensions/attention.py b/transformer_engine/jax/cpp_extensions/attention.py index 489bfde997..e129997b85 100644 --- a/transformer_engine/jax/cpp_extensions/attention.py +++ b/transformer_engine/jax/cpp_extensions/attention.py @@ -75,6 +75,7 @@ "cp_axis", "cp_striped_window_size", "stripe_size", + "return_max_logit", ], ) @dataclass(frozen=True) @@ -99,6 +100,7 @@ class _FusedAttnConfig: stripe_size: ( int | None ) # Only for CP + Striped. For Ring P2P, stripe_size=1 only.For AG, stripe_size>=1. + return_max_logit: bool = False @dataclass(frozen=True) @@ -122,6 +124,7 @@ class FusedAttnHelper: head_dim_qk: int head_dim_v: int window_size: Tuple[int, int] + return_max_logit: bool = False def is_fused_attn_kernel_available(self): """Check if there is available fused attention kernel""" @@ -146,6 +149,7 @@ def get_fused_attn_backend(self): self.head_dim_v, self.window_size[0], self.window_size[1], + self.return_max_logit, not self.is_non_deterministic_allowed(), ) @@ -351,6 +355,7 @@ def abstract( q_head_dim, v_head_dim, config.window_size, + config.return_max_logit, ).get_fused_attn_backend() if backend == NVTE_Fused_Attn_Backend.NVTE_F16_arbitrary_seqlen: @@ -371,6 +376,14 @@ def abstract( else: raise ValueError(f"Unsupported {backend=}") softmax_aux_aval = q_aval.update(shape=softmax_shape, dtype=softmax_dtype) + if config.return_max_logit: + if config.qkv_layout.is_thd() and get_cudnn_version() >= (9, 6, 0): + max_tensor_shape = (*batch_shape, q_max_seqlen, attn_heads, 1) + else: + max_tensor_shape = (*batch_shape, attn_heads, q_max_seqlen, 1) + else: + max_tensor_shape = (0,) + max_tensor_aval = q_aval.update(shape=max_tensor_shape, dtype=softmax_dtype) # JAX does not enable 64-bit int by default so we get XLA to allocate x8 memory with # 32-bit unsigned int to get the buffer size we need in the C++ kernel @@ -417,6 +430,7 @@ def abstract( config.max_segments_per_seq, config.window_size[0], config.window_size[1], + config.return_max_logit, bottom_right_diagonal, ) wkspace_aval = q_aval.update( @@ -437,17 +451,19 @@ def abstract( f" {softmax_offset_aval.shape}" ) - return out_aval, softmax_aux_aval, rng_state_aval, wkspace_aval + return out_aval, softmax_aux_aval, max_tensor_aval, rng_state_aval, wkspace_aval @staticmethod def outer_abstract(*args, **kwargs): """ Fused attention fwd outer primitive abstract """ - out_aval, softmax_aux_aval, rng_state_aval, _ = FusedAttnFwdPrimitive.abstract( + out_aval, softmax_aux_aval, _, rng_state_aval, _ = FusedAttnFwdPrimitive.abstract( *args, **kwargs ) - return out_aval, softmax_aux_aval, rng_state_aval + max_logit_shape = (out_aval.shape[-2],) if kwargs["config"].return_max_logit else (0,) + max_logit_aval = out_aval.update(shape=max_logit_shape, dtype=out_aval.dtype) + return out_aval, softmax_aux_aval, rng_state_aval, max_logit_aval @staticmethod def lowering( @@ -531,6 +547,7 @@ def lowering( mask_type=int(config.attn_mask_type.value), qkv_layout=int(config.qkv_layout.value), is_training=config.is_training, + return_max_logit=config.return_max_logit, deterministic=not FusedAttnHelper.is_non_deterministic_allowed(), window_size_left=window_size_left, window_size_right=window_size_right, @@ -574,6 +591,9 @@ def impl( config.max_segments_per_seq, ) ) + raw_q_seqlen = q_seqlen + raw_q_seq_offsets = q_seq_offsets + if config.qkv_layout.is_thd(): def _fix_len_take(x, condition, fill_value=-1): @@ -630,7 +650,7 @@ def convert_to_2d(offsets, batch, max_seqlen): q_cu_seqlen = generate_cu_seqlen(q_seqlen.flatten()) kv_cu_seqlen = generate_cu_seqlen(kv_seqlen.flatten()) - output, softmax_aux, rng_state, _ = FusedAttnFwdPrimitive.inner_primitive.bind( + output, softmax_aux, max_tensor, rng_state, _ = FusedAttnFwdPrimitive.inner_primitive.bind( q, k, v, @@ -647,7 +667,39 @@ def convert_to_2d(offsets, batch, max_seqlen): _kv_segment_pos, config=config, ) - return output, softmax_aux, rng_state + max_logit = FusedAttnFwdPrimitive._reduce_max_logit( + max_tensor, output, raw_q_seqlen, raw_q_seq_offsets, config + ) + return output, softmax_aux, rng_state, max_logit + + @staticmethod + def _reduce_max_logit(max_tensor, output, q_seqlen, q_seq_offsets, config): + """Reduce cuDNN's raw Max tensor to PyTorch-compatible per-head max_logit.""" + if not config.return_max_logit: + return jnp.zeros((0,), dtype=output.dtype) + + if config.qkv_layout.is_thd() and max_tensor.ndim == 4: + q_seqlen = jnp.where(q_seqlen > 0, q_seqlen, 0) + q_seq_offsets = jnp.where(q_seq_offsets >= 0, q_seq_offsets, -1) + token_idx = jnp.arange(output.shape[-3], dtype=q_seq_offsets.dtype) + valid = jnp.any( + (q_seq_offsets[..., :-1, None] >= 0) + & (token_idx >= q_seq_offsets[..., :-1, None]) + & (token_idx < (q_seq_offsets[..., :-1, None] + q_seqlen[..., None])), + axis=-2, + ) + if max_tensor.shape[1] == output.shape[-3]: + max_tensor = jnp.where(valid[:, :, None, None], max_tensor, -jnp.inf) + else: + max_tensor = jnp.where(valid[:, None, :, None], max_tensor, -jnp.inf) + + if max_tensor.ndim == 3: + amax_dims = (0, 2) + elif config.qkv_layout.is_thd() and max_tensor.shape[1] == output.shape[-3]: + amax_dims = (0, 1, 3) + else: + amax_dims = (0, 2, 3) + return jnp.max(max_tensor, axis=amax_dims).astype(output.dtype) @staticmethod def batcher(batched_args, batch_dims, *, config): @@ -659,7 +711,8 @@ def batcher(batched_args, batch_dims, *, config): q_bdim, _, _, _, _, seed_bdim, *_ = batch_dims # Pass through; segment_ids/segment_pos may have different batch dims (e.g. vmapped ids, # replicated pos). get_seqlens_and_offsets() in attention.py handles conversion without expanding. - out_bdims = q_bdim, q_bdim, seed_bdim + max_logit_bdim = q_bdim if config.return_max_logit else None + out_bdims = q_bdim, q_bdim, seed_bdim, max_logit_bdim return ( FusedAttnFwdPrimitive.outer_primitive.bind(*batched_args, config=config), out_bdims, @@ -713,12 +766,16 @@ def infer_sharding_from_operands(config, mesh, arg_infos, result_infos): raise ValueError(f"Unsupported {config.qkv_layout=}") rng_state_sharding = NamedSharding(mesh, PartitionSpec(get_all_mesh_axes(), None)) - return (out_sharding, softmax_aux_sharding, rng_state_sharding) + max_logit_sharding = NamedSharding( + mesh, PartitionSpec(q_spec[-2] if config.return_max_logit else None) + ) + return (out_sharding, softmax_aux_sharding, rng_state_sharding, max_logit_sharding) @staticmethod def partition(config, mesh, arg_infos, result_infos): out_sharding = result_infos[0].sharding softmax_aux_sharding = result_infos[1].sharding + max_logit_sharding = result_infos[3].sharding rng_state_sharding = seed_sharding = NamedSharding( mesh, PartitionSpec(get_all_mesh_axes(), None) ) @@ -727,7 +784,7 @@ def partition(config, mesh, arg_infos, result_infos): arg_shardings[-1] = arg_shardings[-3] arg_shardings[-2] = arg_shardings[-4] arg_shardings = tuple(arg_shardings) - out_shardings = (out_sharding, softmax_aux_sharding, rng_state_sharding) + out_shardings = (out_sharding, softmax_aux_sharding, rng_state_sharding, max_logit_sharding) impl = partial(FusedAttnFwdPrimitive.impl, config=config) return mesh, impl, out_shardings, arg_shardings @@ -756,8 +813,10 @@ def shardy_sharding_rule(config, mesh, value_types, result_types): else: softmax_aux_sharding = ("…0", "head", "seqlen", "i") + max_logit_sharding = ("head",) if config.return_max_logit else ("max_logit",) return SdyShardingRule( - tuple(input_spec), (out_sharding, softmax_aux_sharding, rng_sharding) + tuple(input_spec), + (out_sharding, softmax_aux_sharding, rng_sharding, max_logit_sharding), ) @@ -1810,19 +1869,22 @@ def partition(config, mesh, arg_infos, result_infos): ), "Sliding window attention is not supported when context parallelism is enabled" if not is_context_parallel: return FusedAttnFwdPrimitive.partition(config, mesh, arg_infos, result_infos) + if config.return_max_logit: + raise NotImplementedError("return_max_logit is not yet supported with context parallelism") helper = _FusedAttnCPWithAllGatherHelper(mesh, config) helper.check_supported() out_sharding = result_infos[0].sharding softmax_aux_sharding = result_infos[1].sharding + max_logit_sharding = result_infos[3].sharding rng_state_sharding = seed_sharding = NamedSharding( mesh, PartitionSpec(get_all_mesh_axes(), None) ) arg_shardings = [arg_i.sharding for arg_i in arg_infos] arg_shardings[5] = seed_sharding arg_shardings = tuple(arg_shardings) - out_shardings = (out_sharding, softmax_aux_sharding, rng_state_sharding) + out_shardings = (out_sharding, softmax_aux_sharding, rng_state_sharding, max_logit_sharding) def impl( q, @@ -1870,7 +1932,7 @@ def _cross_attn(idx, q, k, v, bias, softmax_offset, q_seqlen, kv_seqlen, seed): q_seqlen_for_step = q_seqlen / (cp_size * 2) num_kv_chunks = kv_max_seqlen // kv_seqlens_for_rank[sub_idx] kv_seqlen_for_step = (kv_seqlen / (cp_size * 2)) * num_kv_chunks - output, softmax_aux, rng_state = FusedAttnFwdPrimitive.impl( + output, softmax_aux, rng_state, _ = FusedAttnFwdPrimitive.impl( q_split[sub_idx], k_unmasked, v_unmasked, @@ -1892,8 +1954,9 @@ def _cross_attn(idx, q, k, v, bias, softmax_offset, q_seqlen, kv_seqlen, seed): output = jnp.concatenate((results[0][0], results[1][0]), axis=1) softmax_aux = jnp.concatenate((results[0][1], results[1][1]), axis=2) rng_state = results[1][2] # Use the final RNG state + max_logit = jnp.zeros((0,), dtype=output.dtype) - return output, softmax_aux, rng_state + return output, softmax_aux, rng_state, max_logit k_ag, v_ag = helper.all_gather_kv(k, v) @@ -2103,19 +2166,22 @@ def partition(config, mesh, arg_infos, result_infos): is_context_parallel = get_mesh_axis_size(config.cp_axis, mesh) > 1 if not is_context_parallel: return FusedAttnFwdPrimitive.partition(config, mesh, arg_infos, result_infos) + if config.return_max_logit: + raise NotImplementedError("return_max_logit is not yet supported with context parallelism") helper = _FusedAttnCPWithAllGatherHelper(mesh, config) helper.check_supported() out_sharding = result_infos[0].sharding softmax_aux_sharding = result_infos[1].sharding + max_logit_sharding = result_infos[3].sharding rng_state_sharding = seed_sharding = NamedSharding( mesh, PartitionSpec(get_all_mesh_axes(), None) ) arg_shardings = [arg_i.sharding for arg_i in arg_infos] arg_shardings[5] = seed_sharding arg_shardings = tuple(arg_shardings) - out_shardings = (out_sharding, softmax_aux_sharding, rng_state_sharding) + out_shardings = (out_sharding, softmax_aux_sharding, rng_state_sharding, max_logit_sharding) def impl( q, @@ -2179,7 +2245,7 @@ def _cross_attn( max_segments_per_seq=adjusted_max_segments_per_seq, ) - output, softmax_aux, rng_state = FusedAttnFwdPrimitive.impl( + output, softmax_aux, rng_state, _ = FusedAttnFwdPrimitive.impl( q, # sharded for rank k, # ag v, # ag @@ -2198,7 +2264,8 @@ def _cross_attn( max_seqlen=kv_max_seqlen, cp_size=cp_size ), ) - return output, softmax_aux, rng_state + max_logit = jnp.zeros((0,), dtype=output.dtype) + return output, softmax_aux, rng_state, max_logit # AG the k, v, kv_segment_ids and kv_segment_pos k_ag, v_ag = helper.all_gather_kv(k, v) @@ -2553,12 +2620,15 @@ def partition(config, mesh, arg_infos, result_infos): ), "Sliding window attention is not supported when context parallelism is enabled" if not is_context_parallel: return FusedAttnFwdPrimitive.partition(config, mesh, arg_infos, result_infos) + if config.return_max_logit: + raise NotImplementedError("return_max_logit is not yet supported with context parallelism") helper = _FusedAttnCPWithP2PHelper(mesh, config) helper.check_supported() out_sharding = result_infos[0].sharding softmax_aux_sharding = result_infos[1].sharding + max_logit_sharding = result_infos[3].sharding rng_state_sharding = seed_sharding = NamedSharding( mesh, PartitionSpec(get_all_mesh_axes(), None) ) @@ -2568,7 +2638,7 @@ def partition(config, mesh, arg_infos, result_infos): arg_shardings[-1] = arg_shardings[-3] arg_shardings[-2] = arg_shardings[-4] arg_shardings = tuple(arg_shardings) - out_shardings = (out_sharding, softmax_aux_sharding, rng_state_sharding) + out_shardings = (out_sharding, softmax_aux_sharding, rng_state_sharding, max_logit_sharding) def ring_attn_fwd_impl( q, @@ -2616,7 +2686,7 @@ def scan_kv_block(idx, carry): def mask_compute(attn_mask_type): q_seqlen_per_step = helper.adjust_seqlen(q_seqlen, q_max_seqlen, idx) kv_seqlen_per_step = helper.adjust_seqlen(kv_seqlen, kv_max_seqlen, idx) - output_per_step, softmax_aux_per_step, _ = FusedAttnFwdPrimitive.impl( + output_per_step, softmax_aux_per_step, _, _ = FusedAttnFwdPrimitive.impl( q, kv, _not_used, @@ -2642,7 +2712,7 @@ def half_kv_no_mask_compute(): q_seqlen_per_step = helper.adjust_seqlen(q_seqlen, q_max_seqlen, idx) kv_seqlen_per_step = helper.adjust_seqlen(kv_seqlen, kv_max_seqlen, idx) // 2 kv_part = lax.slice_in_dim(kv, 0, kv.shape[1] // 2, axis=1) - output_per_step, softmax_aux_per_step, _ = FusedAttnFwdPrimitive.impl( + output_per_step, softmax_aux_per_step, _, _ = FusedAttnFwdPrimitive.impl( q, kv_part, _not_used, @@ -2665,7 +2735,7 @@ def half_q_no_mask_compute(): q_seqlen_per_step = helper.adjust_seqlen(q_seqlen, q_max_seqlen, idx) // 2 kv_seqlen_per_step = helper.adjust_seqlen(kv_seqlen, kv_max_seqlen, idx) q_part = lax.slice_in_dim(q, q_max_seqlen // 2, q_max_seqlen, axis=1) - output_per_step, softmax_aux_per_step, _ = FusedAttnFwdPrimitive.impl( + output_per_step, softmax_aux_per_step, _, _ = FusedAttnFwdPrimitive.impl( q_part, kv, _not_used, @@ -2747,7 +2817,8 @@ def correction(output, softmax_aux, output_per_step, softmax_aux_per_step): (kv, output, softmax_aux) = carry output = output.astype(q.dtype) - return output, softmax_aux, rng_state + max_logit = jnp.zeros((0,), dtype=output.dtype) + return output, softmax_aux, rng_state, max_logit return mesh, ring_attn_fwd_impl, out_shardings, arg_shardings @@ -3059,12 +3130,15 @@ def partition(config, mesh, arg_infos, result_infos): is_context_parallel = get_mesh_axis_size(config.cp_axis, mesh) > 1 if not is_context_parallel: return FusedAttnFwdPrimitive.partition(config, mesh, arg_infos, result_infos) + if config.return_max_logit: + raise NotImplementedError("return_max_logit is not yet supported with context parallelism") helper = _FusedAttnCPWithP2PHelper(mesh, config) helper.check_supported() out_sharding = result_infos[0].sharding softmax_aux_sharding = result_infos[1].sharding + max_logit_sharding = result_infos[3].sharding rng_state_sharding = seed_sharding = NamedSharding( mesh, PartitionSpec(get_all_mesh_axes(), None) ) @@ -3074,7 +3148,7 @@ def partition(config, mesh, arg_infos, result_infos): arg_shardings[-1] = arg_shardings[-3] arg_shardings[-2] = arg_shardings[-4] arg_shardings = tuple(arg_shardings) - out_shardings = (out_sharding, softmax_aux_sharding, rng_state_sharding) + out_shardings = (out_sharding, softmax_aux_sharding, rng_state_sharding, max_logit_sharding) def fwd_impl( q, @@ -3157,7 +3231,7 @@ def compute(config): ) else: current_config = subblock_config - output_per_step, softmax_aux_per_step, _ = compute(current_config) + output_per_step, softmax_aux_per_step, _, _ = compute(current_config) softmax_aux_per_step = softmax_aux_per_step.reshape((batch, q_max_seqlen, head, 1)) @@ -3194,7 +3268,9 @@ def correction(output, softmax_aux, output_per_step, softmax_aux_per_step): carry = scan_kv_block(i, carry) (_, _, _, output, softmax_aux) = carry - return output.astype(q.dtype), softmax_aux, rng_state + output = output.astype(q.dtype) + max_logit = jnp.zeros((0,), dtype=output.dtype) + return output, softmax_aux, rng_state, max_logit return mesh, fwd_impl, out_shardings, arg_shardings @@ -3375,6 +3451,7 @@ def fused_attn_fwd( context_parallel_causal_load_balanced: bool = False, context_parallel_axis: str = "", stripe_size: int | None = None, + return_max_logit: bool = False, ) -> jnp.ndarray: """ Perform the forward pass of with cuDNN fused attention implementations. @@ -3414,6 +3491,7 @@ def fused_attn_fwd( Indicates the sequences are ordered for causal mask load balancing when running context parallelism. context_parallel_axis (str): The name of the context parallel axis. stripe_size (int | None): Indicates the striping height to be used for ReorderStrategy.Striped Load Balancing + return_max_logit (bool): Whether to return the per-head maximum attention logit. Returns: (jnp.ndarray): The output tensor from the fused attention. """ @@ -3489,6 +3567,7 @@ def fused_attn_fwd( cp_axis=_maybe_context_parallel_axis(context_parallel_axis), cp_striped_window_size=None, stripe_size=stripe_size, + return_max_logit=return_max_logit, ) primitive = None @@ -3506,7 +3585,7 @@ def fused_attn_fwd( primitive = FusedRingAttnFwdPrimitive.outer_primitive seq_desc_flatten, _ = jax.tree.flatten(sequence_descriptor) - output, softmax_aux, rng_state = primitive.bind( + output, softmax_aux, rng_state, max_logit = primitive.bind( *qkv_for_primitive, bias, softmax_offset, @@ -3515,7 +3594,7 @@ def fused_attn_fwd( config=fused_config, ) rng_state = with_sharding_constraint(rng_state, PartitionSpec(get_all_mesh_axes(), None)) - return (output, softmax_aux, rng_state) + return (output, softmax_aux, rng_state, max_logit) def fused_attn_bwd( From de3b095b95adab646835ef54f171826cca0dbcfc Mon Sep 17 00:00:00 2001 From: Kshitij Lakhani Date: Tue, 9 Jun 2026 23:13:11 -0700 Subject: [PATCH 3/9] Add max logit to JAX fused attn FFI and set it in the workspace Signed-off-by: Kshitij Lakhani --- transformer_engine/jax/csrc/extensions.h | 4 +- .../jax/csrc/extensions/attention.cpp | 80 ++++++++++++------- 2 files changed, 53 insertions(+), 31 deletions(-) diff --git a/transformer_engine/jax/csrc/extensions.h b/transformer_engine/jax/csrc/extensions.h index b9c7c849f2..c7c54d3c64 100644 --- a/transformer_engine/jax/csrc/extensions.h +++ b/transformer_engine/jax/csrc/extensions.h @@ -156,7 +156,7 @@ NVTE_Fused_Attn_Backend GetFusedAttnBackend( NVTE_Bias_Type bias_type, NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, float dropout_probability, size_t q_attn_heads, size_t kv_attn_heads, size_t q_max_seqlen, size_t kv_max_seqlen, size_t qk_head_dim, size_t v_head_dim, int64_t window_size_left, - int64_t window_size_right, bool deterministic); + int64_t window_size_right, bool return_max_logit, bool deterministic); pybind11::tuple GetFusedAttnForwardWorkspaceSizes( size_t input_batch, size_t bias_batch, size_t q_max_seqlen, size_t kv_max_seqlen, @@ -164,7 +164,7 @@ pybind11::tuple GetFusedAttnForwardWorkspaceSizes( size_t v_head_dim, float scaling_factor, float dropout_probability, NVTE_Bias_Type bias_type, NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, NVTE_QKV_Layout qkv_layout, DType dtype, bool is_training, size_t max_segments_per_seq, int64_t window_size_left, - int64_t window_size_right, bool bottom_right_diagonal); + int64_t window_size_right, bool return_max_logit, bool bottom_right_diagonal); pybind11::tuple GetFusedAttnBackwardWorkspaceSizes( size_t input_batch, size_t bias_batch, size_t q_max_seqlen, size_t kv_max_seqlen, diff --git a/transformer_engine/jax/csrc/extensions/attention.cpp b/transformer_engine/jax/csrc/extensions/attention.cpp index 3fd6780d6d..9728f2cbcf 100644 --- a/transformer_engine/jax/csrc/extensions/attention.cpp +++ b/transformer_engine/jax/csrc/extensions/attention.cpp @@ -29,12 +29,12 @@ NVTE_Fused_Attn_Backend GetFusedAttnBackend( NVTE_Bias_Type bias_type, NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, float dropout_probability, size_t q_attn_heads, size_t kv_attn_heads, size_t q_max_seqlen, size_t kv_max_seqlen, size_t qk_head_dim, size_t v_head_dim, int64_t window_size_left, - int64_t window_size_right, bool deterministic) { + int64_t window_size_right, bool return_max_logit, bool deterministic) { auto backend = nvte_get_fused_attn_backend( is_training, static_cast(q_dtype), static_cast(kv_dtype), qkv_layout, bias_type, mask_type, softmax_type, dropout_probability, q_attn_heads, kv_attn_heads, q_max_seqlen, kv_max_seqlen, qk_head_dim, v_head_dim, window_size_left, window_size_right, - false, false, deterministic); + return_max_logit, false, deterministic); return backend; } @@ -48,8 +48,8 @@ void PrepareFusedAttnForwardAuxTensors(NVTETensorPack *tensor_pack, const size_t const size_t bias_heads, const size_t q_max_seqlen, const size_t kv_max_seqlen, DType dtype, NVTE_Bias_Type bias_type, NVTE_Fused_Attn_Backend backend, - void *softmax_buf, void *rng_state_buf = nullptr, - void *bias_buf = nullptr, + void *softmax_buf, void *max_buf = nullptr, + void *rng_state_buf = nullptr, void *bias_buf = nullptr, void *softmax_offset_buf = nullptr) { // all backends need softmax but expect different shapes/dtypes tensor_pack->size = 1; @@ -65,8 +65,23 @@ void PrepareFusedAttnForwardAuxTensors(NVTETensorPack *tensor_pack, const size_t // arbitrary sequence length backend needs the RNG state and a different shape/dtype softmax if (backend == NVTE_Fused_Attn_Backend::NVTE_F16_arbitrary_seqlen) { - tensor_pack->size = 2; - NVTETensor &rng_state_aux = tensor_pack->tensors[1]; + int size = 1; // Start after softmax. + + if (max_buf != nullptr) { + NVTETensor &max_aux = tensor_pack->tensors[size++]; + NVTEBasicTensor max_aux_data; + max_aux_data.data_ptr = max_buf; + max_aux_data.shape = {}; + max_aux_data.shape.ndim = 4; + max_aux_data.shape.data[0] = input_batch; + max_aux_data.shape.data[1] = attn_heads; + max_aux_data.shape.data[2] = q_max_seqlen; + max_aux_data.shape.data[3] = 1; + max_aux_data.dtype = static_cast(DType::kFloat32); + nvte_set_tensor_param(&max_aux, kNVTERowwiseData, &max_aux_data); + } + + NVTETensor &rng_state_aux = tensor_pack->tensors[size++]; NVTEBasicTensor rng_state_aux_data; rng_state_aux_data.data_ptr = rng_state_buf; rng_state_aux_data.shape = {}; @@ -77,8 +92,6 @@ void PrepareFusedAttnForwardAuxTensors(NVTETensorPack *tensor_pack, const size_t softmax_aux_data.shape.data[3] = 1; // {B,H,Qs,Ks} -> {B,H,Qs,1} softmax_aux_data.dtype = static_cast(DType::kFloat32); - int size = 2; // Start at 2 (we have softmax and rng_state at indices 0, 1) - // include bias if enabled if (bias_type != NVTE_Bias_Type::NVTE_NO_BIAS && bias_type != NVTE_Bias_Type::NVTE_ALIBI) { NVTETensor &bias_aux = tensor_pack->tensors[size]; @@ -136,7 +149,7 @@ void PrepareFusedAttnBackwardAuxTensors(NVTETensorPack *tensor_pack, const size_ auto dummy_backend = NVTE_Fused_Attn_Backend::NVTE_F16_arbitrary_seqlen; PrepareFusedAttnForwardAuxTensors(tensor_pack, input_batch, bias_batch, attn_heads, bias_heads, q_max_seqlen, kv_max_seqlen, dtype, dummy_bias_type, - dummy_backend, softmax_buf, rng_state_buf, bias_buf, + dummy_backend, softmax_buf, nullptr, rng_state_buf, bias_buf, softmax_offset_buf); } @@ -146,7 +159,7 @@ pybind11::tuple GetFusedAttnForwardWorkspaceSizes( size_t v_head_dim, float scaling_factor, float dropout_probability, NVTE_Bias_Type bias_type, NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, NVTE_QKV_Layout qkv_layout, DType dtype, bool is_training, size_t max_segments_per_seq, int64_t window_size_left, - int64_t window_size_right, bool bottom_right_diagonal) { + int64_t window_size_right, bool return_max_logit, bool bottom_right_diagonal) { auto is_ragged = nvte_get_qkv_format(qkv_layout) == NVTE_QKV_Format::NVTE_THD; auto q_shape = is_ragged ? std::vector{input_batch * q_max_seqlen, attn_heads, qk_head_dim} @@ -200,8 +213,8 @@ pybind11::tuple GetFusedAttnForwardWorkspaceSizes( dummy_softmax_offset_tensor.data(), s_tensor.data(), o_tensor.data(), &aux_output_tensors, q_cu_seqlens_tensor.data(), kv_cu_seqlens_tensor.data(), ragged_offset_tensor.data(), ragged_offset_tensor.data(), dummy_page_table_tensor.data(), dummy_page_table_tensor.data(), - dummy_rng_state_tensor.data(), q_max_seqlen, kv_max_seqlen, is_training, false, false, - scaling_factor, dropout_probability, qkv_layout, nvte_get_q_format(qkv_layout), + dummy_rng_state_tensor.data(), q_max_seqlen, kv_max_seqlen, is_training, return_max_logit, + false, scaling_factor, dropout_probability, qkv_layout, nvte_get_q_format(qkv_layout), NVTE_QKV_Format_NOT_SET, bias_type, mask_type, softmax_type, window_size_left, window_size_right, bottom_right_diagonal, query_workspace_tensor.data(), nullptr); } @@ -242,13 +255,14 @@ pybind11::tuple GetFusedAttnForwardWorkspaceSizes( static void FusedAttnForwardImpl( cudaStream_t stream, void *q, void *k, void *v, void *bias, void *softmax_offset, void *seed, void *q_cu_seqlens, void *kv_cu_seqlens, void *q_seq_offsets, void *k_seq_offsets, void *output, - void *softmax_aux, void *rng_state, void *workspace, size_t input_batch, size_t bias_batch, - size_t q_max_seqlen, size_t kv_max_seqlen, size_t attn_heads, size_t num_gqa_groups, - size_t bias_heads, size_t qk_head_dim, size_t v_head_dim, size_t max_segments_per_seq, - size_t wkspace_size, float scaling_factor, float dropout_probability, NVTE_Bias_Type bias_type, - NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, NVTE_QKV_Layout qkv_layout, - DType dtype, DType wkspace_dtype, bool is_training, bool deterministic, - int64_t window_size_left, int64_t window_size_right, bool bottom_right_diagonal) { + void *softmax_aux, void *max_tensor, void *rng_state, void *workspace, size_t input_batch, + size_t bias_batch, size_t q_max_seqlen, size_t kv_max_seqlen, size_t attn_heads, + size_t num_gqa_groups, size_t bias_heads, size_t qk_head_dim, size_t v_head_dim, + size_t max_segments_per_seq, size_t wkspace_size, float scaling_factor, + float dropout_probability, NVTE_Bias_Type bias_type, NVTE_Mask_Type mask_type, + NVTE_Softmax_Type softmax_type, NVTE_QKV_Layout qkv_layout, DType dtype, DType wkspace_dtype, + bool is_training, bool return_max_logit, bool deterministic, int64_t window_size_left, + int64_t window_size_right, bool bottom_right_diagonal) { FUSED_ATTN_IMPL_COMMON_BLOCK; /* Input tensors */ @@ -263,6 +277,9 @@ static void FusedAttnForwardImpl( // Memset to 0xF0 for filling large negative numbers auto softmax_aux_size = input_batch * q_max_seqlen * attn_heads; cudaMemsetAsync(softmax_aux, 0xF0, softmax_aux_size * sizeof(float), stream); + if (return_max_logit) { + cudaMemsetAsync(max_tensor, 0xF0, softmax_aux_size * sizeof(float), stream); + } } /* Output tensors */ @@ -278,7 +295,7 @@ static void FusedAttnForwardImpl( is_training, static_cast(dtype), static_cast(dtype), qkv_layout, bias_type, mask_type, softmax_type, dropout_probability, attn_heads, num_gqa_groups, q_max_seqlen, kv_max_seqlen, qk_head_dim, v_head_dim, window_size_left, window_size_right, - false, false, deterministic); + return_max_logit, false, deterministic); nvte_populate_rng_state_async(rng_state, seed, q_max_seqlen, kv_max_seqlen, backend, stream); /* Auxiliary tensors (to be propagated to the backward pass later) */ @@ -286,7 +303,8 @@ static void FusedAttnForwardImpl( nvte_tensor_pack_create(&aux_output_tensors); PrepareFusedAttnForwardAuxTensors(&aux_output_tensors, input_batch, bias_batch, attn_heads, bias_heads, q_max_seqlen, kv_max_seqlen, dtype, bias_type, - backend, softmax_aux, softmax_offset); + backend, softmax_aux, return_max_logit ? max_tensor : nullptr, + rng_state, bias, softmax_offset); /* Call the underlying NVTE API */ auto dummy_page_table_tensor = TensorWrapper(nullptr, std::vector{1}, DType::kInt32); @@ -344,7 +362,7 @@ static void FusedAttnForwardImpl( softmax_offset_tensor.data(), s_tensor.data(), o_tensor.data(), &aux_output_tensors, q_cu_seqlens_tensor.data(), kv_cu_seqlens_tensor.data(), q_seq_offsets_tensor.data(), k_seq_offsets_tensor.data(), dummy_page_table_tensor.data(), dummy_page_table_tensor.data(), - rng_state_tensor.data(), q_max_seqlen, kv_max_seqlen, is_training, false, false, + rng_state_tensor.data(), q_max_seqlen, kv_max_seqlen, is_training, return_max_logit, false, scaling_factor, dropout_probability, qkv_layout, nvte_get_q_format(qkv_layout), NVTE_QKV_Format_NOT_SET, bias_type, mask_type, softmax_type, window_size_left, window_size_right, bottom_right_diagonal, workspace_tensor.data(), stream); @@ -378,6 +396,7 @@ static void FusedAttnForwardImpl( NVTE_QKV_Layout qkv_layout = \ static_cast(get_attr_value(attrs, "qkv_layout")); \ bool is_training = get_attr_value(attrs, "is_training"); \ + bool return_max_logit = get_attr_value_or_default(attrs, "return_max_logit", false); \ bool deterministic = get_attr_value(attrs, "deterministic"); \ auto is_ragged = nvte_get_qkv_format(qkv_layout) == NVTE_QKV_Format::NVTE_THD; \ size_t wkspace_size = product(workspace_buf->dimensions()); \ @@ -390,8 +409,9 @@ Error_Type FusedAttnForwardFFI(cudaStream_t stream, Buffer_Type q_buf, Buffer_Ty Buffer_Type q_cu_seqlens_buf, Buffer_Type kv_cu_seqlens_buf, Buffer_Type q_seq_offsets_buf, Buffer_Type k_seq_offsets_buf, Variadic_Buffer_Type _unused_args, Result_Type output_buf, - Result_Type softmax_aux_buf, Result_Type rng_state_buf, - Result_Type workspace_buf, Dictionary attrs) { + Result_Type softmax_aux_buf, Result_Type max_tensor_buf, + Result_Type rng_state_buf, Result_Type workspace_buf, + Dictionary attrs) { FUSED_ATTN_FFI_GET_ATTRS; FusedAttnForwardImpl( @@ -400,11 +420,12 @@ Error_Type FusedAttnForwardFFI(cudaStream_t stream, Buffer_Type q_buf, Buffer_Ty q_cu_seqlens_buf.untyped_data(), kv_cu_seqlens_buf.untyped_data(), is_ragged ? q_seq_offsets_buf.untyped_data() : nullptr, is_ragged ? k_seq_offsets_buf.untyped_data() : nullptr, output_buf->untyped_data(), - softmax_aux_buf->untyped_data(), rng_state_buf->untyped_data(), workspace_buf->untyped_data(), - input_batch, bias_batch, q_max_seqlen, kv_max_seqlen, attn_heads, num_gqa_groups, bias_heads, - qk_head_dim, v_head_dim, max_segments_per_seq, wkspace_size, scaling_factor, - dropout_probability, bias_type, mask_type, softmax_type, qkv_layout, dtype, wkspace_dtype, - is_training, deterministic, window_size_left, window_size_right, bottom_right_diagonal); + softmax_aux_buf->untyped_data(), max_tensor_buf->untyped_data(), + rng_state_buf->untyped_data(), workspace_buf->untyped_data(), input_batch, bias_batch, + q_max_seqlen, kv_max_seqlen, attn_heads, num_gqa_groups, bias_heads, qk_head_dim, v_head_dim, + max_segments_per_seq, wkspace_size, scaling_factor, dropout_probability, bias_type, mask_type, + softmax_type, qkv_layout, dtype, wkspace_dtype, is_training, return_max_logit, deterministic, + window_size_left, window_size_right, bottom_right_diagonal); return ffi_with_cuda_error_check(); } @@ -424,6 +445,7 @@ XLA_FFI_DEFINE_HANDLER_SYMBOL(FusedAttnForwardHandler, FusedAttnForwardFFI, .RemainingArgs() // _cp_aux_args unused .Ret() // output .Ret() // softmax_aux + .Ret() // max_tensor .Ret() // rng_state .Ret() // workspace .Attrs(), From 99d52d4f1b2e83604e5fae374063cbf66287001f Mon Sep 17 00:00:00 2001 From: Kshitij Lakhani Date: Tue, 9 Jun 2026 23:15:21 -0700 Subject: [PATCH 4/9] Add first pass tests for max logit and softmax aux tensor outputs in JAX fused attn tests Signed-off-by: Kshitij Lakhani --- tests/jax/test_fused_attn.py | 224 +++++++++++++++++++++++++++++++++-- 1 file changed, 216 insertions(+), 8 deletions(-) diff --git a/tests/jax/test_fused_attn.py b/tests/jax/test_fused_attn.py index 0d1db0b9e1..532a8dacd0 100644 --- a/tests/jax/test_fused_attn.py +++ b/tests/jax/test_fused_attn.py @@ -79,27 +79,23 @@ def general_dot_product_attention( dropout_rng: ArrayLike, dtype: DTypeLike, score_mod_reference: Optional[Callable[[Array], Array]] = None, + is_max_logit_enabled: bool = False, ) -> Array: """ Similar to flax.linen.dot_product_attention but with GQA support """ query, key, value, bias = promote_dtype(query, key, value, bias, dtype=dtype) dtype = query.dtype - b, s_q, h_q, d = query.shape _, s_kv, h_kv, _ = key.shape assert (h_q % h_kv == 0) and (h_q >= h_kv) num_groups = h_q // h_kv grouped_query = jnp.reshape(query, (b, s_q, h_kv, num_groups, d)) - # logits with shape (b, h_kv, num_groups, s_q, s_kv) logits = scale_factor * jnp.einsum("...qhgd,...khd->...hgqk", grouped_query, key) if bias is not None: - # reshape logits without groups logits = logits.reshape((b, h_kv * num_groups, s_q, s_kv)) - # apply post-scale bias logits = logits + bias - # reshape logits back to original logits = logits.reshape((b, h_kv, num_groups, s_q, s_kv)) if mask is not None: @@ -110,6 +106,8 @@ def general_dot_product_attention( if score_mod_reference is not None: # Kernel tests use NO_MASK; fused_attn rejects mask+score_mod before this reference path. logits = score_mod_reference(logits.astype(jnp.float32)) + if is_max_logit_enabled: + return jnp.max(logits.reshape((b, h_q, s_q, s_kv)), axis=(0, 2, 3)).astype(dtype) match softmax_type: case AttnSoftmaxType.VANILLA_SOFTMAX: @@ -268,7 +266,17 @@ def _split_valid_and_invalid(primitive, reference, pad): return primitive_valid, primitive_invalid, reference_valid, reference_invalid -def jax_dpa(query, key, value, bias, softmax_offset, mask, dropout_rng, **kwargs): +def jax_dpa( + query, + key, + value, + bias, + softmax_offset, + mask, + dropout_rng, + is_max_logit_enabled=False, + **kwargs, +): """ JAX native dot product attention implementation """ @@ -308,6 +316,7 @@ def jax_dpa(query, key, value, bias, softmax_offset, mask, dropout_rng, **kwargs dropout_rng=dropout_rng, dtype=jnp.float32, score_mod_reference=score_mod_reference, + is_max_logit_enabled=is_max_logit_enabled, ) return output.astype(query.dtype) @@ -339,9 +348,13 @@ def customcall_fused_dpa( qkv_args = (query, key, value) case _: raise ValueError(f"Unsupported {qkv_layout=}") - return fused_attn( + result = fused_attn( qkv_args, bias, sequence_descriptor, dropout_rng, softmax_offset=softmax_offset, **kwargs - ).astype(query.dtype) + ) + if isinstance(result, tuple): + output, aux = result + return output.astype(query.dtype), aux + return result.astype(query.dtype) def test_fused_attn_score_mod_rejects_masks_before_cudnn_frontend(): @@ -1227,6 +1240,201 @@ def check_dqkv(primitive, reference, pad, idx): target_hlo = jitted_primitive.lower(*customcall_args).compile().as_text() assert_equal_collectives(target_hlo, self.coll_count_ref) + def _reference_args(self): + return [self.q, self.k, self.v, self.bias, self.softmax_offset, self.mask, self.dropout_rng] + + def _customcall_args(self): + return [ + jax.device_put(self.cp_reorder_fn(self.q), self.qkvo_sharding), + jax.device_put(self.cp_reorder_fn(self.k), self.qkvo_sharding), + jax.device_put(self.cp_reorder_fn(self.v), self.qkvo_sharding), + jax.device_put(self.bias, self.bias_sharding), + jax.device_put(self.softmax_offset, self.softmax_offset_sharding), + jax.device_put(self.sequence_desciptor, self.seq_desc_sharding), + jax.device_put(self.dropout_rng, self.dropout_rng_sharding), + ] + + def _fused_attn_kwargs(self, **overrides): + kwargs = { + "attn_bias_type": self.attn_bias_type, + "attn_mask_type": self.attn_mask_type, + "softmax_type": self.softmax_type, + "scaling_factor": self.scaling_factor, + "dropout_probability": self.dropout_prob, + "is_training": self.is_training, + "qkv_layout": self.qkv_layout, + "max_segments_per_seq": self._get_max_segments_per_sequence(), + "window_size": self.window_size, + "context_parallel_strategy": self.cp_strategy, + "context_parallel_causal_load_balanced": self.cp_load_balanced, + "stripe_size": self.stripe_size, + } + kwargs.update(overrides) + return kwargs + + def test_forward_with_max_logit(self): + """Test forward output and returned max_logit.""" + self._setup_inputs() + kwargs = self._fused_attn_kwargs() + + customcall_fused_dpa_jit = jit( + partial(customcall_fused_dpa, return_max_logit=True, **kwargs), + static_argnames=kwargs.keys(), + in_shardings=[ + self.qkvo_sharding, + self.qkvo_sharding, + self.qkvo_sharding, + self.bias_sharding, + self.softmax_offset_sharding, + self.seq_desc_sharding, + self.dropout_rng_sharding, + ], + ) + + with self.mesh, autocast(mesh_resource=self.mesh_resource): + primitive_out, primitive_aux = customcall_fused_dpa_jit(*self._customcall_args()) + primitive_max_logit = primitive_aux["max_logit"] + primitive_out = self.cp_inverse_reorder_fn(primitive_out) + + reference_out = jax_dpa(*self._reference_args(), **kwargs) + reference_max_logit = jax_dpa( + *self._reference_args(), is_max_logit_enabled=True, **kwargs + ) + + primitive_valid, primitive_invalid, reference_valid, _ = _split_valid_and_invalid( + primitive_out, reference_out, self.pad_q + ) + assert_allclose(primitive_invalid, jnp.zeros_like(primitive_invalid), dtype=self.dtype) + assert_allclose(primitive_valid, reference_valid, dtype=self.dtype) + assert_allclose(primitive_max_logit, reference_max_logit, dtype=self.dtype) + + def test_forward_with_softmax_aux(self): + """Test optional softmax_aux return wiring.""" + self._setup_inputs() + kwargs = self._fused_attn_kwargs(return_softmax_aux=True) + + customcall_fused_dpa_jit = jit( + partial(customcall_fused_dpa, **kwargs), + static_argnames=kwargs.keys(), + in_shardings=[ + self.qkvo_sharding, + self.qkvo_sharding, + self.qkvo_sharding, + self.bias_sharding, + self.softmax_offset_sharding, + self.seq_desc_sharding, + self.dropout_rng_sharding, + ], + ) + + with self.mesh, autocast(mesh_resource=self.mesh_resource): + output, aux = customcall_fused_dpa_jit(*self._customcall_args()) + + assert output.shape == self.q.shape + assert "softmax_aux" in aux + assert aux["softmax_aux"].dtype == jnp.float32 + + def test_backward_with_max_logit(self): + """Ensure aux-return cotangents do not break the fused attention backward path.""" + self._setup_inputs() + kwargs = self._fused_attn_kwargs() + + def loss_fn(query): + output, _ = customcall_fused_dpa( + query, + self.k, + self.v, + self.bias, + self.softmax_offset, + self.sequence_desciptor, + self.dropout_rng, + return_max_logit=True, + **kwargs, + ) + return jnp.mean(output.astype(jnp.float32)) + + grad = jax.grad(loss_fn)(self.q) + assert grad.shape == self.q.shape + + +@pytest.mark.parametrize( + "qkv_layout, seq_desc_format", + [ + pytest.param(QKVLayout.BSHD_BSHD_BSHD, SeqDescFormat.Seqlens, id="BSHD_SEPARATE"), + pytest.param(QKVLayout.T3HD, SeqDescFormat.Seqlens, id="THD_QKV_PACKED"), + ], +) +def test_fused_attn_return_max_logit(qkv_layout, seq_desc_format): + """Check non-CP JAX fused attention can expose PyTorch-compatible max_logit.""" + runner = FusedAttnRunner( + batch_size=2, + max_seqlen_q=128, + max_seqlen_kv=128, + num_heads_q=8, + num_heads_kv=8, + head_dim_qk=64, + head_dim_v=64, + attn_bias_type=AttnBiasType.NO_BIAS, + attn_mask_type=AttnMaskType.PADDING_CAUSAL_MASK, + softmax_type=AttnSoftmaxType.VANILLA_SOFTMAX, + dropout_prob=0.0, + dtype=jnp.bfloat16, + is_training=True, + qkv_layout=qkv_layout, + bias_shape=None, + window_size=None, + seq_desc_format=seq_desc_format, + ) + runner.test_forward_with_max_logit() + + +def test_fused_attn_return_softmax_aux(): + """Check the optional public softmax_aux return is wired for non-CP attention.""" + runner = FusedAttnRunner( + batch_size=2, + max_seqlen_q=128, + max_seqlen_kv=128, + num_heads_q=8, + num_heads_kv=8, + head_dim_qk=64, + head_dim_v=64, + attn_bias_type=AttnBiasType.NO_BIAS, + attn_mask_type=AttnMaskType.PADDING_CAUSAL_MASK, + softmax_type=AttnSoftmaxType.VANILLA_SOFTMAX, + dropout_prob=0.0, + dtype=jnp.bfloat16, + is_training=True, + qkv_layout=QKVLayout.BS3HD, + bias_shape=None, + window_size=None, + seq_desc_format=SeqDescFormat.Seqlens, + ) + runner.test_forward_with_softmax_aux() + + +def test_fused_attn_return_max_logit_backward_smoke(): + """Ensure aux-return cotangents do not break the fused attention backward path.""" + runner = FusedAttnRunner( + batch_size=2, + max_seqlen_q=128, + max_seqlen_kv=128, + num_heads_q=8, + num_heads_kv=8, + head_dim_qk=64, + head_dim_v=64, + attn_bias_type=AttnBiasType.NO_BIAS, + attn_mask_type=AttnMaskType.PADDING_CAUSAL_MASK, + softmax_type=AttnSoftmaxType.VANILLA_SOFTMAX, + dropout_prob=0.0, + dtype=jnp.bfloat16, + is_training=True, + qkv_layout=QKVLayout.BSHD_BSHD_BSHD, + bias_shape=None, + window_size=None, + seq_desc_format=SeqDescFormat.Seqlens, + ) + runner.test_backward_with_max_logit() + def _get_swa_window_size_for_test(s_kv: int, attn_mask_type: AttnMaskType) -> Tuple[int, int]: """Pick a sliding-window size for SWA tests, gated on cuDNN version. From e55e230c6780650fffe2465c558a54830a0758ff Mon Sep 17 00:00:00 2001 From: Kshitij Lakhani Date: Wed, 8 Jul 2026 16:00:37 -0700 Subject: [PATCH 5/9] Reject aux returns with score_mod Signed-off-by: Kshitij Lakhani --- transformer_engine/jax/attention.py | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/transformer_engine/jax/attention.py b/transformer_engine/jax/attention.py index 9d3aaae90d..a300fb353c 100644 --- a/transformer_engine/jax/attention.py +++ b/transformer_engine/jax/attention.py @@ -1629,6 +1629,18 @@ def fused_attn( if score_mod_only_args: raise ValueError(f"{', '.join(score_mod_only_args)} require score_mod to be provided.") else: + aux_return_args = [ + name + for name, value in ( + ("return_max_logit", return_max_logit), + ("return_softmax_aux", return_softmax_aux), + ) + if value + ] + if aux_return_args: + raise ValueError( + f"{', '.join(aux_return_args)} are not supported with score_mod fused_attn." + ) tex.validate_fused_attn_score_mod( qkv, bias, From 47ac14331dfe5afdf957566de2e63fe484ac3d8b Mon Sep 17 00:00:00 2001 From: Kshitij Lakhani Date: Wed, 8 Jul 2026 16:01:43 -0700 Subject: [PATCH 6/9] Handle SM120 max-logit layout Signed-off-by: Kshitij Lakhani --- .../jax/cpp_extensions/attention.py | 16 +++++++++++++--- 1 file changed, 13 insertions(+), 3 deletions(-) diff --git a/transformer_engine/jax/cpp_extensions/attention.py b/transformer_engine/jax/cpp_extensions/attention.py index e129997b85..53db542798 100644 --- a/transformer_engine/jax/cpp_extensions/attention.py +++ b/transformer_engine/jax/cpp_extensions/attention.py @@ -377,7 +377,7 @@ def abstract( raise ValueError(f"Unsupported {backend=}") softmax_aux_aval = q_aval.update(shape=softmax_shape, dtype=softmax_dtype) if config.return_max_logit: - if config.qkv_layout.is_thd() and get_cudnn_version() >= (9, 6, 0): + if FusedAttnFwdPrimitive._uses_thd_ragged_max_tensor(config): max_tensor_shape = (*batch_shape, q_max_seqlen, attn_heads, 1) else: max_tensor_shape = (*batch_shape, attn_heads, q_max_seqlen, 1) @@ -678,6 +678,7 @@ def _reduce_max_logit(max_tensor, output, q_seqlen, q_seq_offsets, config): if not config.return_max_logit: return jnp.zeros((0,), dtype=output.dtype) + uses_thd_ragged_max_tensor = FusedAttnFwdPrimitive._uses_thd_ragged_max_tensor(config) if config.qkv_layout.is_thd() and max_tensor.ndim == 4: q_seqlen = jnp.where(q_seqlen > 0, q_seqlen, 0) q_seq_offsets = jnp.where(q_seq_offsets >= 0, q_seq_offsets, -1) @@ -688,19 +689,28 @@ def _reduce_max_logit(max_tensor, output, q_seqlen, q_seq_offsets, config): & (token_idx < (q_seq_offsets[..., :-1, None] + q_seqlen[..., None])), axis=-2, ) - if max_tensor.shape[1] == output.shape[-3]: + if uses_thd_ragged_max_tensor: max_tensor = jnp.where(valid[:, :, None, None], max_tensor, -jnp.inf) else: max_tensor = jnp.where(valid[:, None, :, None], max_tensor, -jnp.inf) if max_tensor.ndim == 3: amax_dims = (0, 2) - elif config.qkv_layout.is_thd() and max_tensor.shape[1] == output.shape[-3]: + elif uses_thd_ragged_max_tensor: amax_dims = (0, 1, 3) else: amax_dims = (0, 2, 3) return jnp.max(max_tensor, axis=amax_dims).astype(output.dtype) + @staticmethod + def _uses_thd_ragged_max_tensor(config): + """Return whether cuDNN writes THD Max with BSH-like ragged-stats layout.""" + return ( + config.qkv_layout.is_thd() + and get_cudnn_version() >= (9, 6, 0) + and 120 not in get_all_device_compute_capability() + ) + @staticmethod def batcher(batched_args, batch_dims, *, config): # batch_dims: each element is the batch axis (0, ...) or None. Only 0 or None allowed. From ef6ba3f5086840cf78ed12818b66e72d2378abd1 Mon Sep 17 00:00:00 2001 From: Kshitij Lakhani Date: Wed, 8 Jul 2026 16:02:12 -0700 Subject: [PATCH 7/9] Drop softmax aux return Signed-off-by: Kshitij Lakhani --- tests/jax/test_fused_attn.py | 50 ------------------------ transformer_engine/jax/attention.py | 60 +++++++---------------------- 2 files changed, 14 insertions(+), 96 deletions(-) diff --git a/tests/jax/test_fused_attn.py b/tests/jax/test_fused_attn.py index 532a8dacd0..bcf536196c 100644 --- a/tests/jax/test_fused_attn.py +++ b/tests/jax/test_fused_attn.py @@ -1308,32 +1308,6 @@ def test_forward_with_max_logit(self): assert_allclose(primitive_valid, reference_valid, dtype=self.dtype) assert_allclose(primitive_max_logit, reference_max_logit, dtype=self.dtype) - def test_forward_with_softmax_aux(self): - """Test optional softmax_aux return wiring.""" - self._setup_inputs() - kwargs = self._fused_attn_kwargs(return_softmax_aux=True) - - customcall_fused_dpa_jit = jit( - partial(customcall_fused_dpa, **kwargs), - static_argnames=kwargs.keys(), - in_shardings=[ - self.qkvo_sharding, - self.qkvo_sharding, - self.qkvo_sharding, - self.bias_sharding, - self.softmax_offset_sharding, - self.seq_desc_sharding, - self.dropout_rng_sharding, - ], - ) - - with self.mesh, autocast(mesh_resource=self.mesh_resource): - output, aux = customcall_fused_dpa_jit(*self._customcall_args()) - - assert output.shape == self.q.shape - assert "softmax_aux" in aux - assert aux["softmax_aux"].dtype == jnp.float32 - def test_backward_with_max_logit(self): """Ensure aux-return cotangents do not break the fused attention backward path.""" self._setup_inputs() @@ -1388,30 +1362,6 @@ def test_fused_attn_return_max_logit(qkv_layout, seq_desc_format): runner.test_forward_with_max_logit() -def test_fused_attn_return_softmax_aux(): - """Check the optional public softmax_aux return is wired for non-CP attention.""" - runner = FusedAttnRunner( - batch_size=2, - max_seqlen_q=128, - max_seqlen_kv=128, - num_heads_q=8, - num_heads_kv=8, - head_dim_qk=64, - head_dim_v=64, - attn_bias_type=AttnBiasType.NO_BIAS, - attn_mask_type=AttnMaskType.PADDING_CAUSAL_MASK, - softmax_type=AttnSoftmaxType.VANILLA_SOFTMAX, - dropout_prob=0.0, - dtype=jnp.bfloat16, - is_training=True, - qkv_layout=QKVLayout.BS3HD, - bias_shape=None, - window_size=None, - seq_desc_format=SeqDescFormat.Seqlens, - ) - runner.test_forward_with_softmax_aux() - - def test_fused_attn_return_max_logit_backward_smoke(): """Ensure aux-return cotangents do not break the fused attention backward path.""" runner = FusedAttnRunner( diff --git a/transformer_engine/jax/attention.py b/transformer_engine/jax/attention.py index a300fb353c..3eb41ae13f 100644 --- a/transformer_engine/jax/attention.py +++ b/transformer_engine/jax/attention.py @@ -1056,7 +1056,6 @@ def _legacy_fused_attn( context_parallel_axis: str = "", softmax_offset: Optional[jnp.ndarray] = None, return_max_logit: bool = False, - return_softmax_aux: bool = False, ): """ Perform non-THD (non-packed) cuDNN fused attention. @@ -1092,14 +1091,12 @@ def _legacy_fused_attn( [1, num_heads, 1, 1]. Used when softmax_type is AttnSoftmaxType.LEARNABLE_SOFTMAX. return_max_logit (bool): If True, also return per-head maximum attention logits in an auxiliary dictionary under ``"max_logit"``. - return_softmax_aux (bool): If True, also return backend-specific softmax statistics - in an auxiliary dictionary under ``"softmax_aux"``. Returns: jnp.ndarray: - Attention output when neither ``return_max_logit`` nor ``return_softmax_aux`` is True. + Attention output when ``return_max_logit`` is False. tuple[jnp.ndarray, dict[str, jnp.ndarray]]: - ``(output, aux)`` when either flag is True. ``aux`` may contain: - ``"max_logit"`` (shape ``[h]``) and/or ``"softmax_aux"`` (float32). + ``(output, aux)`` when ``return_max_logit`` is True. ``aux`` contains + ``"max_logit"`` with shape ``[h]``. """ assert ( not qkv_layout.is_thd() @@ -1154,7 +1151,6 @@ def _legacy_fused_attn( context_parallel_causal_load_balanced=context_parallel_causal_load_balanced, context_parallel_axis=context_parallel_axis, return_max_logit=return_max_logit, - return_softmax_aux=return_softmax_aux, ) return output @@ -1181,7 +1177,6 @@ def fused_attn_thd( context_parallel_axis: str = "", softmax_offset: Optional[jnp.ndarray] = None, return_max_logit: bool = False, - return_softmax_aux: bool = False, ): """ Deprecated THD fused attn, please use fusd_attn with SequenceDescriptor @@ -1237,7 +1232,6 @@ def fused_attn_thd( context_parallel_causal_load_balanced=context_parallel_causal_load_balanced, context_parallel_axis=context_parallel_axis, return_max_logit=return_max_logit, - return_softmax_aux=return_softmax_aux, ) return output @@ -1245,7 +1239,7 @@ def fused_attn_thd( @partial( jax.custom_vjp, - nondiff_argnums=(5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20), + nondiff_argnums=(5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19), ) def _fused_attn( qkv: Tuple[jnp.ndarray, ...], @@ -1268,7 +1262,6 @@ def _fused_attn( context_checkpoint_name: str = "context", stripe_size: int | None = None, return_max_logit: bool = False, - return_softmax_aux: bool = False, ): output, _ = _fused_attn_fwd_rule( qkv, @@ -1291,7 +1284,6 @@ def _fused_attn( context_checkpoint_name=context_checkpoint_name, stripe_size=stripe_size, return_max_logit=return_max_logit, - return_softmax_aux=return_softmax_aux, ) return output @@ -1317,7 +1309,6 @@ def _fused_attn_fwd_rule( context_checkpoint_name, stripe_size, return_max_logit, - return_softmax_aux, ): output, softmax_aux, rng_state, max_logit = tex.fused_attn_fwd( qkv, @@ -1344,9 +1335,7 @@ def _fused_attn_fwd_rule( softmax_aux = checkpoint_name(softmax_aux, context_checkpoint_name) rng_state = checkpoint_name(rng_state, context_checkpoint_name) max_logit = checkpoint_name(max_logit, context_checkpoint_name) - attn_output = _resolve_fused_attn_output( - output, max_logit, softmax_aux, return_max_logit, return_softmax_aux - ) + attn_output = _resolve_fused_attn_output(output, max_logit, return_max_logit) return attn_output, ( qkv, bias, @@ -1374,12 +1363,11 @@ def _fused_attn_bwd_rule( context_checkpoint_name, stripe_size, return_max_logit, - return_softmax_aux, ctx, dz, ): del context_checkpoint_name - if return_max_logit or return_softmax_aux: + if return_max_logit: dz, _ = dz ( qkv, @@ -1426,16 +1414,11 @@ def _fused_attn_bwd_rule( ) -def _resolve_fused_attn_output(output, max_logit, softmax_aux, return_max_logit, return_softmax_aux): - if not return_max_logit and not return_softmax_aux: +def _resolve_fused_attn_output(output, max_logit, return_max_logit): + if not return_max_logit: return output - aux = {} - if return_max_logit: - aux["max_logit"] = max_logit - if return_softmax_aux: - aux["softmax_aux"] = softmax_aux - return output, aux + return output, {"max_logit": max_logit} _fused_attn.defvjp(_fused_attn_fwd_rule, _fused_attn_bwd_rule) @@ -1519,7 +1502,6 @@ def fused_attn( score_mod_tensors: Optional[Mapping[str, Any]] = None, score_mod_bprop_tensors: Optional[Mapping[str, Any]] = None, return_max_logit: bool = False, - return_softmax_aux: bool = False, ): """ Perform cuDNN fused attention. @@ -1578,14 +1560,12 @@ def fused_attn( Python/NumPy scalars made available to `score_mod_bprop`. return_max_logit (bool): If True, also return per-head maximum attention logits in an auxiliary dictionary under ``"max_logit"``. - return_softmax_aux (bool): If True, also return backend-specific softmax statistics - in an auxiliary dictionary under ``"softmax_aux"``. Returns: jnp.ndarray: - Attention output when neither ``return_max_logit`` nor ``return_softmax_aux`` is True. + Attention output when ``return_max_logit`` is False. tuple[jnp.ndarray, dict[str, jnp.ndarray]]: - ``(output, aux)`` when either flag is True. ``aux`` may contain: - ``"max_logit"`` (shape ``[h]``) and/or ``"softmax_aux"`` (float32). + ``(output, aux)`` when ``return_max_logit`` is True. ``aux`` contains + ``"max_logit"`` with shape ``[h]``. Examples (non-THD, also known as non-packed): >>> # q_segment_ids = [[1, 1, 1, 0], [1, 1, 0, 0]], 0 means padded tokens @@ -1629,18 +1609,8 @@ def fused_attn( if score_mod_only_args: raise ValueError(f"{', '.join(score_mod_only_args)} require score_mod to be provided.") else: - aux_return_args = [ - name - for name, value in ( - ("return_max_logit", return_max_logit), - ("return_softmax_aux", return_softmax_aux), - ) - if value - ] - if aux_return_args: - raise ValueError( - f"{', '.join(aux_return_args)} are not supported with score_mod fused_attn." - ) + if return_max_logit: + raise ValueError("return_max_logit is not supported with score_mod fused_attn.") tex.validate_fused_attn_score_mod( qkv, bias, @@ -1701,7 +1671,6 @@ def fused_attn( context_parallel_axis=context_parallel_axis, softmax_offset=softmax_offset, return_max_logit=return_max_logit, - return_softmax_aux=return_softmax_aux, ) if max_segments_per_seq > 1 and not qkv_layout.is_thd(): warnings.warn( @@ -1733,6 +1702,5 @@ def fused_attn( context_checkpoint_name=context_checkpoint_name, stripe_size=stripe_size, return_max_logit=return_max_logit, - return_softmax_aux=return_softmax_aux, ) return output From c3f8d35498b08205cd4fe8fa8a1e802d0e5e3aad Mon Sep 17 00:00:00 2001 From: Kshitij Lakhani Date: Mon, 20 Jul 2026 10:29:53 -0700 Subject: [PATCH 8/9] Modify static args in fused attn tests for jax Signed-off-by: Kshitij Lakhani --- tests/jax/test_fused_attn.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/jax/test_fused_attn.py b/tests/jax/test_fused_attn.py index bcf536196c..f64ecbe91f 100644 --- a/tests/jax/test_fused_attn.py +++ b/tests/jax/test_fused_attn.py @@ -64,7 +64,7 @@ def init(): yield -@partial(jax.jit, static_argnums=(6, 7, 8, 9, 11, 12)) +@partial(jax.jit, static_argnums=(6, 7, 8, 9, 11, 12, 13)) def general_dot_product_attention( query: ArrayLike, key: ArrayLike, From fc151f99842fab3ed8d7f32cf3bcba124a42f631 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Wed, 22 Jul 2026 01:16:37 +0000 Subject: [PATCH 9/9] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- tests/jax/test_fused_attn.py | 4 +--- .../jax/cpp_extensions/attention.py | 16 ++++++++++++---- 2 files changed, 13 insertions(+), 7 deletions(-) diff --git a/tests/jax/test_fused_attn.py b/tests/jax/test_fused_attn.py index f64ecbe91f..b1990619c6 100644 --- a/tests/jax/test_fused_attn.py +++ b/tests/jax/test_fused_attn.py @@ -1297,9 +1297,7 @@ def test_forward_with_max_logit(self): primitive_out = self.cp_inverse_reorder_fn(primitive_out) reference_out = jax_dpa(*self._reference_args(), **kwargs) - reference_max_logit = jax_dpa( - *self._reference_args(), is_max_logit_enabled=True, **kwargs - ) + reference_max_logit = jax_dpa(*self._reference_args(), is_max_logit_enabled=True, **kwargs) primitive_valid, primitive_invalid, reference_valid, _ = _split_valid_and_invalid( primitive_out, reference_out, self.pad_q diff --git a/transformer_engine/jax/cpp_extensions/attention.py b/transformer_engine/jax/cpp_extensions/attention.py index 53db542798..487199620c 100644 --- a/transformer_engine/jax/cpp_extensions/attention.py +++ b/transformer_engine/jax/cpp_extensions/attention.py @@ -1880,7 +1880,9 @@ def partition(config, mesh, arg_infos, result_infos): if not is_context_parallel: return FusedAttnFwdPrimitive.partition(config, mesh, arg_infos, result_infos) if config.return_max_logit: - raise NotImplementedError("return_max_logit is not yet supported with context parallelism") + raise NotImplementedError( + "return_max_logit is not yet supported with context parallelism" + ) helper = _FusedAttnCPWithAllGatherHelper(mesh, config) helper.check_supported() @@ -2177,7 +2179,9 @@ def partition(config, mesh, arg_infos, result_infos): if not is_context_parallel: return FusedAttnFwdPrimitive.partition(config, mesh, arg_infos, result_infos) if config.return_max_logit: - raise NotImplementedError("return_max_logit is not yet supported with context parallelism") + raise NotImplementedError( + "return_max_logit is not yet supported with context parallelism" + ) helper = _FusedAttnCPWithAllGatherHelper(mesh, config) helper.check_supported() @@ -2631,7 +2635,9 @@ def partition(config, mesh, arg_infos, result_infos): if not is_context_parallel: return FusedAttnFwdPrimitive.partition(config, mesh, arg_infos, result_infos) if config.return_max_logit: - raise NotImplementedError("return_max_logit is not yet supported with context parallelism") + raise NotImplementedError( + "return_max_logit is not yet supported with context parallelism" + ) helper = _FusedAttnCPWithP2PHelper(mesh, config) helper.check_supported() @@ -3141,7 +3147,9 @@ def partition(config, mesh, arg_infos, result_infos): if not is_context_parallel: return FusedAttnFwdPrimitive.partition(config, mesh, arg_infos, result_infos) if config.return_max_logit: - raise NotImplementedError("return_max_logit is not yet supported with context parallelism") + raise NotImplementedError( + "return_max_logit is not yet supported with context parallelism" + ) helper = _FusedAttnCPWithP2PHelper(mesh, config) helper.check_supported()