From ed56839bbe792c839d020f98e5adde5c01ee1299 Mon Sep 17 00:00:00 2001 From: Kamil Nasyrov Date: Thu, 9 Jul 2026 19:14:01 +0000 Subject: [PATCH 1/9] add FWD pass for atten kernels with sink from CK --- .../include/ck_fused_attn/ck_fused_attn.hpp | 5 + .../ck_fused_attn/src/ck_fused_attn_fwd.cpp | 9 +- .../common/fused_attn_rocm/fused_attn.cpp | 5 +- .../common/fused_attn_rocm/fused_attn_ck.cpp | 110 ++++++++++-------- .../common/fused_attn_rocm/fused_attn_ck.h | 4 +- .../jax/csrc/extensions/attention.cpp | 2 - 6 files changed, 76 insertions(+), 59 deletions(-) diff --git a/transformer_engine/common/ck_fused_attn/include/ck_fused_attn/ck_fused_attn.hpp b/transformer_engine/common/ck_fused_attn/include/ck_fused_attn/ck_fused_attn.hpp index 127d75b4ca..d85ae75b15 100644 --- a/transformer_engine/common/ck_fused_attn/include/ck_fused_attn/ck_fused_attn.hpp +++ b/transformer_engine/common/ck_fused_attn/include/ck_fused_attn/ck_fused_attn.hpp @@ -79,6 +79,11 @@ struct CKAttnCommonArgs { void* philox_seed_ptr = nullptr; void* philox_offset_ptr = nullptr; + // Softmax sink (learnable / off-by-one) + const void* sink_ptr = nullptr; + bool has_sink = false; + int64_t sink_size = 0; + // O layout (o_ptr lives in derived because fwd writes it / bwd reads it) uint64_t stride_b_o = 0, stride_h_o = 0, stride_s_o = 0; diff --git a/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp b/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp index 0f407230c0..c59a2e6b37 100644 --- a/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp +++ b/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp @@ -91,6 +91,9 @@ void log_fwd_config(const char* func_name, bool has_dropout, const aiter::mha_fw log_value(log_file, "dropout_seed_ptr", std::get<0>(std::get>(fmha_args.drop_seed_offset))); log_value(log_file, "dropout_offset_ptr", std::get<1>(std::get>(fmha_args.drop_seed_offset))); + log_value(log_file, "has_sink", fmha_args.has_sink); + log_value(log_file, "sink_ptr", fmha_args.sink_ptr); + log_value(log_file, "sink_size", fmha_args.sink_size); } void dump_fwd_timings(const char* dump_path, float average_runtime){ @@ -160,7 +163,7 @@ hipError_t ck_attn_fwd(const CKAttnFwdArgs& args, hipStream_t stream){ fmha_args.block_scale_seqstart_q_ptr = nullptr; fmha_args.block_scale_seqstart_k_ptr = nullptr; - fmha_args.sink_ptr = nullptr; + fmha_args.sink_ptr = args.sink_ptr; fmha_args.seqlen_k = args.s_kv; // unused in group mode (or kvcache enabled) fmha_args.max_seqlen_q = args.s_q; @@ -207,11 +210,11 @@ hipError_t ck_attn_fwd(const CKAttnFwdArgs& args, hipStream_t stream){ fmha_args.bias_type = static_cast(bias_type); fmha_args.has_lse = args.lse_ptr!=nullptr; fmha_args.qscale_type = static_cast(quant_scale_enum::no_scale); - fmha_args.has_sink = false; + fmha_args.has_sink = args.has_sink; fmha_args.q_descale_ptr = nullptr; fmha_args.k_descale_ptr = nullptr; fmha_args.v_descale_ptr = nullptr; - fmha_args.sink_size = 0; + fmha_args.sink_size = args.sink_size; fmha_args.min_seqlen_q = 0; fmha_args.block_scale_size_q = 0; fmha_args.block_scale_size_kv = 0; diff --git a/transformer_engine/common/fused_attn_rocm/fused_attn.cpp b/transformer_engine/common/fused_attn_rocm/fused_attn.cpp index 943fa16972..4834f0ef02 100644 --- a/transformer_engine/common/fused_attn_rocm/fused_attn.cpp +++ b/transformer_engine/common/fused_attn_rocm/fused_attn.cpp @@ -382,6 +382,7 @@ void nvte_fused_attn_fwd(const NVTETensor Q, const NVTETensor K, const NVTETenso const Tensor *input_K = convertNVTETensorCheck(K); const Tensor *input_V = convertNVTETensorCheck(V); const Tensor *input_Bias = convertNVTETensorCheck(Bias); + const Tensor *input_SoftmaxOffset = convertNVTETensorCheck(SoftmaxOffset); Tensor *output_O = convertNVTETensorCheck(O); Tensor *wkspace = convertNVTETensorCheck(workspace); @@ -413,9 +414,9 @@ void nvte_fused_attn_fwd(const NVTETensor Q, const NVTETensor K, const NVTETenso fused_attn_ck_fwd( b, h_q, h_kv, max_seqlen_q, max_seqlen_kv, d_qk, d_v, is_training, attn_scale, dropout, - qkv_layout, bias_type, attn_mask_type, + qkv_layout, bias_type, attn_mask_type, softmax_type, window_size_left, window_size_right, - input_Q, input_K, input_V, input_Bias, + input_Q, input_K, input_V, input_Bias, input_SoftmaxOffset, output_O, Aux_CTX_Tensors, input_cu_seqlens_q, input_cu_seqlens_kv, diff --git a/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp b/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp index 597387555f..89e056866d 100644 --- a/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp +++ b/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp @@ -81,14 +81,6 @@ bool is_ck_backend_supported( return false; } - // filter based on softmax type - if(softmax_type!=NVTE_VANILLA_SOFTMAX){ - if(nvte_log_ck_config){ - std::cout<<"AITER/CK fused attn does not support learnable sink yet"<data.shape[1]; } + void *devPtrSoftmaxOffset = nullptr; + if (softmax_type != NVTE_VANILLA_SOFTMAX) { + devPtrSoftmaxOffset = input_SoftmaxOffset->data.dptr; + } + void *devPtrCuSeqlensQ = input_cu_seqlens_q->data.dptr; void *devPtrCuSeqlensKV = input_cu_seqlens_kv->data.dptr; void *devPtrSeqOffsetsQ = input_cu_seqlens_q_padded->data.dptr; @@ -1150,53 +1159,51 @@ void fused_attn_ck_fwd( size_t max_tokens_q = std::accumulate((input_Q->data).shape.begin(), (input_Q->data).shape.end(), static_cast(1), std::multiplies())/h_q/d_qk; size_t max_tokens_kv = std::accumulate((input_K->data).shape.begin(), (input_K->data).shape.end(), static_cast(1), std::multiplies())/h_kv/d_qk; - bool is_ragged = nvte_get_qkv_format(qkv_layout)==NVTE_QKV_Format::NVTE_THD; + bool is_ragged = nvte_get_qkv_format(qkv_layout)==NVTE_QKV_Format::NVTE_THD; + size_t i = 0; if (Aux_CTX_Tensors->size == 0) { + Tensor *output_S = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); + output_S->data.dptr = nullptr; + if(is_ragged){ + output_S->data.shape = {max_tokens_q, h_q, 1}; + }else{ + output_S->data.shape = {b, h_q, max_seqlen_q, 1}; + } + output_S->data.dtype = DType::kFloat32; + + Tensor *output_rng_state = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); + output_rng_state->data.dptr = nullptr; + output_rng_state->data.shape = {2}; + output_rng_state->data.dtype = DType::kInt64; + if ((bias_type != NVTE_NO_BIAS) && (bias_type != NVTE_ALIBI)) { - Aux_CTX_Tensors->size = 3; - Tensor *output_S = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[0]); - output_S->data.dptr = nullptr; - if(is_ragged){ - output_S->data.shape = {max_tokens_q, h_q, 1}; - }else{ - output_S->data.shape = {b, h_q, max_seqlen_q, 1}; - } - output_S->data.dtype = DType::kFloat32; - Tensor *output_rng_state = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[1]); - output_rng_state->data.dptr = nullptr; - output_rng_state->data.shape = {2}; - output_rng_state->data.dtype = DType::kInt64; - Tensor *output_bias = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[2]); + Tensor *output_bias = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); output_bias->data.dptr = nullptr; output_bias->data.shape = {bias_b, bias_h, max_seqlen_q, max_seqlen_kv}; output_bias->data.dtype = QKV_type; - } else { - Aux_CTX_Tensors->size = 2; - Tensor *output_S = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[0]); - output_S->data.dptr = nullptr; - if(is_ragged){ - output_S->data.shape = {max_tokens_q, h_q, 1}; - }else{ - output_S->data.shape = {b, h_q, max_seqlen_q, 1}; - } - output_S->data.dtype = DType::kFloat32; - Tensor *output_rng_state = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[1]); - output_rng_state->data.dptr = nullptr; - output_rng_state->data.shape = {2}; - output_rng_state->data.dtype = DType::kInt64; } - } else if (Aux_CTX_Tensors->size == 2) { - Tensor *output_S = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[0]); - devPtrS = output_S->data.dptr; - Tensor *output_rng_state = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[1]); - output_rng_state->data.dptr = rng_state->data.dptr; - } else if (Aux_CTX_Tensors->size == 3) { - Tensor *output_S = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[0]); + + if (softmax_type != NVTE_VANILLA_SOFTMAX) { + Tensor *output_softmax_offset = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); + output_softmax_offset->data.dptr = nullptr; + output_softmax_offset->data.shape = {1, h_q, 1, 1}; + output_softmax_offset->data.dtype = DType::kFloat32; + } + + Aux_CTX_Tensors->size = i; + } else if (Aux_CTX_Tensors->size >= 2) { + Tensor *output_S = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); devPtrS = output_S->data.dptr; - Tensor *output_rng_state = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[1]); + Tensor *output_rng_state = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); output_rng_state->data.dptr = rng_state->data.dptr; - Tensor *output_bias = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[2]); - output_bias->data.dptr = devPtrBias; + if ((bias_type != NVTE_NO_BIAS) && (bias_type != NVTE_ALIBI)) { + Tensor *output_bias = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); + output_bias->data.dptr = devPtrBias; + } + if (softmax_type != NVTE_VANILLA_SOFTMAX) { + Tensor *output_softmax_offset = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); + output_softmax_offset->data.dptr = devPtrSoftmaxOffset; + } } else { NVTE_ERROR("Unexpected Aux_CTX_Tensors->size."); } @@ -1211,9 +1218,10 @@ void fused_attn_ck_fwd( max_tokens_q, max_tokens_kv, is_training, attn_scale, dropout, qkv_layout, - bias_type, attn_mask_type, + bias_type, attn_mask_type, softmax_type, window_size_left, window_size_right, - devPtrQ, devPtrK, devPtrV, devPtrBias, + devPtrQ, devPtrK, devPtrV, devPtrBias, + devPtrSoftmaxOffset, devPtrS, devPtrO, rng_state->data.dptr, reinterpret_cast(reinterpret_cast(rng_state->data.dptr) + 1), diff --git a/transformer_engine/common/fused_attn_rocm/fused_attn_ck.h b/transformer_engine/common/fused_attn_rocm/fused_attn_ck.h index 7aafc883fa..e5ebd34720 100644 --- a/transformer_engine/common/fused_attn_rocm/fused_attn_ck.h +++ b/transformer_engine/common/fused_attn_rocm/fused_attn_ck.h @@ -36,8 +36,10 @@ void fused_attn_ck_fwd( size_t b, size_t h_q, size_t h_kv, size_t max_seqlen_q, size_t max_seqlen_kv, size_t d_qk, size_t d_v, bool is_training, float attn_scale, float dropout, NVTE_QKV_Layout qkv_layout, NVTE_Bias_Type bias_type, NVTE_Mask_Type attn_mask_type, + NVTE_Softmax_Type softmax_type, int64_t window_size_left, int64_t window_size_right, - const Tensor* input_Q, const Tensor* input_K, const Tensor* input_V, const Tensor* input_Bias, + const Tensor* input_Q, const Tensor* input_K, const Tensor* input_V, const Tensor* input_Bias, + const Tensor* input_SoftmaxOffset, Tensor* output_O, NVTETensorPack *Aux_CTX_Tensors, const Tensor* input_cu_seqlens_q, const Tensor* input_cu_seqlens_kv, diff --git a/transformer_engine/jax/csrc/extensions/attention.cpp b/transformer_engine/jax/csrc/extensions/attention.cpp index 02efd1b388..fbf9229fa1 100644 --- a/transformer_engine/jax/csrc/extensions/attention.cpp +++ b/transformer_engine/jax/csrc/extensions/attention.cpp @@ -91,7 +91,6 @@ void PrepareFusedAttnForwardAuxTensors(NVTETensorPack *tensor_pack, const size_t bias_aux_data.dtype = static_cast(dtype); nvte_set_tensor_param(&bias_aux, kNVTERowwiseData, &bias_aux_data); } -#ifndef USE_ROCM // include softmax_offset if provided if (softmax_offset_buf != nullptr) { NVTETensor &softmax_offset_aux = tensor_pack->tensors[size]; @@ -106,7 +105,6 @@ void PrepareFusedAttnForwardAuxTensors(NVTETensorPack *tensor_pack, const size_t softmax_offset_aux_data.dtype = static_cast(DType::kFloat32); nvte_set_tensor_param(&softmax_offset_aux, kNVTERowwiseData, &softmax_offset_aux_data); } -#endif // Set final size tensor_pack->size = size; From cd725205453d75323fec11e3bacfeb1290017d35 Mon Sep 17 00:00:00 2001 From: shurale-nkn Date: Fri, 24 Jul 2026 14:57:23 +0000 Subject: [PATCH 2/9] add fmha-sink bwd support --- tests/jax/test_fused_attn.py | 61 ++++++++++++++----- .../include/ck_fused_attn/ck_fused_attn.hpp | 3 + .../ck_fused_attn/src/ck_fused_attn_bwd.cpp | 4 ++ .../common/fused_attn_rocm/fused_attn.cpp | 12 +++- .../common/fused_attn_rocm/fused_attn_ck.cpp | 60 +++++++++++++++++- .../common/fused_attn_rocm/fused_attn_ck.h | 3 + .../jax/csrc/extensions/attention.cpp | 22 ++++--- 7 files changed, 139 insertions(+), 26 deletions(-) diff --git a/tests/jax/test_fused_attn.py b/tests/jax/test_fused_attn.py index 78fb2f1037..1c34d262ed 100644 --- a/tests/jax/test_fused_attn.py +++ b/tests/jax/test_fused_attn.py @@ -996,18 +996,26 @@ def grad_func(func, *args, cp_reverse_out=False, **kwargs): "stripe_size": self.stripe_size, } + arg_nums = (0, 1, 2) + grad_shardings = (self.qkvo_sharding, self.qkvo_sharding, self.qkvo_sharding) + + optional_dgrad_idx = 3 + # We can compute dBias only for the [1, h, s, s] layout - if self.bias_shape == BiasShape._1HSS: - arg_nums = (0, 1, 2, 3) - grad_shardings = ( - self.qkvo_sharding, - self.qkvo_sharding, - self.qkvo_sharding, - self.bias_sharding, - ) - else: - arg_nums = (0, 1, 2) - grad_shardings = (self.qkvo_sharding, self.qkvo_sharding, self.qkvo_sharding) + compute_dbias = self.bias_shape == BiasShape._1HSS + if compute_dbias: + arg_nums += (3,) + grad_shardings += (self.bias_sharding,) + dgrad_idx_dbias = optional_dgrad_idx + optional_dgrad_idx += 1 + + # dsoftmax_offset is only meaningful for the learnable softmax variant + compute_dsoftmax_offset = self.softmax_type == AttnSoftmaxType.LEARNABLE_SOFTMAX + if compute_dsoftmax_offset: + arg_nums += (4,) + grad_shardings += (self.softmax_offset_sharding,) + dgrad_idx_dsoftmax_offset = optional_dgrad_idx + optional_dgrad_idx += 1 # Use FP16/BF16 to sum the results may cause overflow, use FP32 for the summation jitted_primitive = jit( @@ -1102,11 +1110,11 @@ def check_dqkv(primitive, reference, pad, idx): check_dqkv(primitive_dk, reference_dk, self.pad_kv, 1) check_dqkv(primitive_dv, reference_dv, self.pad_kv, 2) - if self.attn_bias_type != AttnBiasType.NO_BIAS and self.bias_shape == BiasShape._1HSS: + if self.attn_bias_type != AttnBiasType.NO_BIAS and compute_dbias: # TODO(mgoldfarb-nvidia): Inverse reorder bias once supported by a CP implementation. - primitive_dbias = primitive_dgrad[3] - reference_dbias = reference_dgrad[3] + primitive_dbias = primitive_dgrad[dgrad_idx_dbias] + reference_dbias = reference_dgrad[dgrad_idx_dbias] # Assume all batch has the same actual_seqlen, probably needs to extend the tests bias_mask = self.mask[0, 0] @@ -1132,6 +1140,31 @@ def check_dqkv(primitive, reference, pad, idx): dtype=self.dtype, ) + if compute_dsoftmax_offset: + primitive_dsoftmax_offset = primitive_dgrad[dgrad_idx_dsoftmax_offset] + reference_dsoftmax_offset = reference_dgrad[dgrad_idx_dsoftmax_offset] + + print_debug_tensor_stats("primitive_dsoftmax_offset", primitive_dsoftmax_offset) + print_debug_tensor_stats("reference_dsoftmax_offset", reference_dsoftmax_offset) + print_debug_tensor_stats( + "diff_dsoftmax_offset", + jnp.abs(primitive_dsoftmax_offset - reference_dsoftmax_offset), + ) + + if is_hip_extension(): + assert not jnp.any( + jnp.isnan(primitive_dsoftmax_offset) + ), "Fused dsoftmax_offset contains NaN" + assert not jnp.any( + jnp.isinf(primitive_dsoftmax_offset) + ), "Fused dsoftmax_offset contains Inf" + + assert_allclose( + primitive_dsoftmax_offset, + reference_dsoftmax_offset, + dtype=self.softmax_offset.dtype, + ) + if self.coll_count_ref is not None: with jax.set_mesh(self.mesh), autocast(mesh_resource=self.mesh_resource): target_hlo = jitted_primitive.lower(*customcall_args).compile().as_text() diff --git a/transformer_engine/common/ck_fused_attn/include/ck_fused_attn/ck_fused_attn.hpp b/transformer_engine/common/ck_fused_attn/include/ck_fused_attn/ck_fused_attn.hpp index d85ae75b15..013d4b6d55 100644 --- a/transformer_engine/common/ck_fused_attn/include/ck_fused_attn/ck_fused_attn.hpp +++ b/transformer_engine/common/ck_fused_attn/include/ck_fused_attn/ck_fused_attn.hpp @@ -136,6 +136,9 @@ struct CkAttnBwdArgs : CKAttnCommonArgs { void* dbias_expanded_ptr = nullptr; void* dbias_ptr = nullptr; + // Softmax sink gradient + void* d_sink_ptr = nullptr; + // Workspace shared with forward LSE void* lse_workspace_ptr = nullptr; diff --git a/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_bwd.cpp b/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_bwd.cpp index b0ed9ae6e5..ddb7082d1a 100644 --- a/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_bwd.cpp +++ b/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_bwd.cpp @@ -366,6 +366,8 @@ void log_bwd_config(const char* func_name, const aiter::mha_bwd_args& fmha_args) log_value(log_file, "dv_ptr", fmha_args.dv_ptr); log_value(log_file, "dbias_ptr", fmha_args.dbias_ptr); log_value(log_file, "dq_acc_ptr", fmha_args.dq_acc_ptr); + log_value(log_file, "sink_ptr", fmha_args.sink_ptr); + log_value(log_file, "d_sink_ptr", fmha_args.d_sink_ptr); log_value(log_file, "seqstart_q_ptr", fmha_args.seqstart_q_ptr); log_value(log_file, "seqstart_k_ptr", fmha_args.seqstart_k_ptr); @@ -496,6 +498,8 @@ hipError_t ck_attn_bwd(const CkAttnBwdArgs& args, hipStream_t stream){ ? (bias_shape==BiasShape::kBHSS ? args.dbias_ptr : args.dbias_expanded_ptr) : nullptr; fmha_args.dq_acc_ptr = args.dq_acc_ptr; + fmha_args.sink_ptr = args.sink_ptr; + fmha_args.d_sink_ptr = args.d_sink_ptr; if (args.is_group_mode()) { fmha_args.seqstart_q_ptr = args.cu_seqlen_q_padded_ptr==nullptr? args.cu_seqlen_q_ptr : args.cu_seqlen_q_padded_ptr; diff --git a/transformer_engine/common/fused_attn_rocm/fused_attn.cpp b/transformer_engine/common/fused_attn_rocm/fused_attn.cpp index 4fcc2d6791..c40e097857 100644 --- a/transformer_engine/common/fused_attn_rocm/fused_attn.cpp +++ b/transformer_engine/common/fused_attn_rocm/fused_attn.cpp @@ -461,6 +461,8 @@ void nvte_fused_attn_bwd(const NVTETensor Q, const NVTETensor K, const NVTETenso const Tensor *output_S = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[0]); //softmax lse const Tensor *input_rng_state = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[1]); Tensor *input_Bias = nullptr; + Tensor *input_SoftmaxOffset = nullptr; + Tensor *output_dSoftmaxOffset = convertNVTETensorCheck(dSoftmaxOffset); auto ndim = input_Q->data.shape.size(); size_t b = input_cu_seqlens_q->data.shape[0] - 1; @@ -487,18 +489,26 @@ void nvte_fused_attn_bwd(const NVTETensor Q, const NVTETensor K, const NVTETenso cuda_graph, deterministic); if (fused_attention_backend == NVTE_Fused_Attn_Backend::NVTE_CK) { + size_t ctx_next_id = 2; if ((bias_type != NVTE_NO_BIAS) && (bias_type != NVTE_ALIBI)) { - input_Bias = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[2]); + input_Bias = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[ctx_next_id++]); + } + if (softmax_type != NVTE_VANILLA_SOFTMAX) { + input_SoftmaxOffset = + convertNVTETensorCheck(Aux_CTX_Tensors->tensors[ctx_next_id++]); } fused_attn_ck_bwd( b, h_q, h_kv, max_seqlen_q, max_seqlen_kv, d_qk, d_v, attn_scale, dropout, qkv_layout, bias_type, attn_mask_type, + softmax_type, window_size_left, window_size_right, deterministic, input_Q, input_K, input_V, input_O, input_dO, input_Bias, + input_SoftmaxOffset, output_S, output_dQ, output_dK, output_dV, output_dBias, + output_dSoftmaxOffset, input_cu_seqlens_q, input_cu_seqlens_kv, input_cu_seqlens_q_padded, diff --git a/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp b/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp index b6256cc75d..8e76d8dca2 100644 --- a/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp +++ b/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp @@ -77,6 +77,16 @@ bool is_ck_backend_supported( return false; } + // WA for failed validation in CK FMHA with SWA and non-vanilla sink + const bool is_swa = window_size_left >= 0; + if (is_swa && softmax_type != NVTE_VANILLA_SOFTMAX) { + if (nvte_log_ck_config) { + std::cout << "CK fused attn does not support non-vanilla sink with SWA" + << std::endl; + } + return false; + } + // joint filter based on sliding window and attn_mask bool is_causal = is_causal_mask(attn_mask_type); if(is_causal){ @@ -137,6 +147,16 @@ bool is_ck_backend_supported( } return false; } + + // WA failed fwd test + bool is_no_mask_window = (window_size_left == -1 && window_size_right == -1); + if(is_ragged && !is_causal && is_no_mask_window && (max_seqlen_q != max_seqlen_kv)){ + if(nvte_log_ck_config){ + std::cout<<"Ragged (THD) cross-attention with plain no-mask padding and " + "max_seqlen_q != max_seqlen_kv is not supported by CK yet"<data.shape[1]; } + void *devPtrSoftmaxOffset = nullptr; + void *devPtrDSoftmaxOffset = nullptr; + if (softmax_type != NVTE_VANILLA_SOFTMAX) { + devPtrSoftmaxOffset = input_SoftmaxOffset->data.dptr; + devPtrDSoftmaxOffset = output_dSoftmaxOffset->data.dptr; + } + void *devPtrdQ = output_dQ->data.dptr; void *devPtrdK = output_dK->data.dptr; void *devPtrdV = output_dV->data.dptr; @@ -1273,12 +1324,15 @@ void fused_attn_ck_bwd( attn_scale, dropout, qkv_layout, bias_type, attn_mask_type, + softmax_type, window_size_left, window_size_right, deterministic, devPtrQ, devPtrK, devPtrV, devPtrO, devPtrSoftmaxStats, devPtrBias, + devPtrSoftmaxOffset, devPtrdQ, devPtrdK, devPtrdV, devPtrdO, devPtrdBias, + devPtrDSoftmaxOffset, rng_state->data.dptr, reinterpret_cast(reinterpret_cast(rng_state->data.dptr) + 1), devPtrCuSeqlensQ, devPtrCuSeqlensKV, diff --git a/transformer_engine/common/fused_attn_rocm/fused_attn_ck.h b/transformer_engine/common/fused_attn_rocm/fused_attn_ck.h index e5ebd34720..1ec34b69ab 100644 --- a/transformer_engine/common/fused_attn_rocm/fused_attn_ck.h +++ b/transformer_engine/common/fused_attn_rocm/fused_attn_ck.h @@ -53,12 +53,15 @@ void fused_attn_ck_bwd( size_t b, size_t h_q, size_t h_kv, size_t max_seqlen_q, size_t max_seqlen_kv, size_t d_qk, size_t d_v, float attn_scale, float dropout, NVTE_QKV_Layout qkv_layout, NVTE_Bias_Type bias_type, NVTE_Mask_Type attn_mask_type, + NVTE_Softmax_Type softmax_type, int64_t window_size_left, int64_t window_size_right, bool deterministic, const Tensor* input_Q, const Tensor* input_K, const Tensor* input_V, const Tensor* input_O, const Tensor* input_dO, const Tensor* input_Bias, + const Tensor* input_SoftmaxOffset, const Tensor* output_S, Tensor* output_dQ, Tensor* output_dK, Tensor* output_dV, Tensor* output_dBias, + Tensor* output_dSoftmaxOffset, const Tensor* input_cu_seqlens_q, const Tensor* input_cu_seqlens_kv, const Tensor* input_cu_seqlens_q_padded, diff --git a/transformer_engine/jax/csrc/extensions/attention.cpp b/transformer_engine/jax/csrc/extensions/attention.cpp index fbf9229fa1..9edf0df01c 100644 --- a/transformer_engine/jax/csrc/extensions/attention.cpp +++ b/transformer_engine/jax/csrc/extensions/attention.cpp @@ -122,22 +122,28 @@ void PrepareFusedAttnForwardAuxTensors(NVTETensorPack *tensor_pack, const size_t void PrepareFusedAttnBackwardAuxTensors(NVTETensorPack *tensor_pack, const size_t input_batch, const size_t bias_batch, const size_t attn_heads, const size_t bias_heads, const size_t q_max_seqlen, - const size_t kv_max_seqlen, DType dtype, - NVTE_Fused_Attn_Backend backend, void *softmax_buf, + const size_t kv_max_seqlen, DType dtype, + [[maybe_unused]]NVTE_Bias_Type bias_type, + [[maybe_unused]]NVTE_Fused_Attn_Backend backend, + void *softmax_buf, void *rng_state_buf, void *bias_buf, void *softmax_offset_buf = nullptr) { +#ifndef USE_ROCM // Backward calls put everything into the tensor pack for every backend // so we set dummy bias_type and backend choices here to follow the correct code path auto dummy_bias_type = NVTE_Bias_Type::NVTE_POST_SCALE_BIAS; -#ifndef USE_ROCM - auto dummy_backend = NVTE_Fused_Attn_Backend::NVTE_F16_arbitrary_seqlen; -#else auto dummy_backend = NVTE_Fused_Attn_Backend::NVTE_No_Backend; -#endif 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, softmax_offset_buf); +#else + PrepareFusedAttnForwardAuxTensors(tensor_pack, input_batch, bias_batch, attn_heads, bias_heads, + q_max_seqlen, kv_max_seqlen, dtype, bias_type, + backend, softmax_buf, rng_state_buf, bias_buf, + softmax_offset_buf); +#endif + #ifndef USE_ROCM // correct softmax shape for max512 sequence length kernel @@ -543,8 +549,8 @@ static void FusedAttnBackwardImpl( q_max_seqlen, kv_max_seqlen, qk_head_dim, v_head_dim, window_size_left, window_size_right, false, false, deterministic); PrepareFusedAttnBackwardAuxTensors(&aux_input_tensors, input_batch, bias_batch, attn_heads, - bias_heads, q_max_seqlen, kv_max_seqlen, dtype, backend, - softmax_aux, rng_state, bias, softmax_offset); + bias_heads, q_max_seqlen, kv_max_seqlen, dtype, bias_type, + backend, softmax_aux, rng_state, bias, softmax_offset); /* Call the underly NVTE API */ // Prepare Q, K, V pointers and shapes based on layout From 6f4a2fcfc724e47093cd6a2e726dcf0fa01dc7d5 Mon Sep 17 00:00:00 2001 From: shurale-nkn Date: Fri, 24 Jul 2026 19:11:46 +0000 Subject: [PATCH 3/9] Update submodule --- 3rdparty/QoLA | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/3rdparty/QoLA b/3rdparty/QoLA index bbcf79e610..ed52843ab1 160000 --- a/3rdparty/QoLA +++ b/3rdparty/QoLA @@ -1 +1 @@ -Subproject commit bbcf79e610c69b4d84615c0bb51abfee7bad1a8d +Subproject commit ed52843ab193a1d8fe63a6a4ebf2198fb0af934d From 4268e24227cfec292ce9191ee016b654d640da45 Mon Sep 17 00:00:00 2001 From: shurale-nkn Date: Sun, 26 Jul 2026 13:53:16 +0000 Subject: [PATCH 4/9] post merge fix --- 3rdparty/QoLA | 2 +- .../common/ck_fused_attn/src/ck_fused_attn_bwd.cpp | 6 ++---- 2 files changed, 3 insertions(+), 5 deletions(-) diff --git a/3rdparty/QoLA b/3rdparty/QoLA index ed52843ab1..d48c1d4ed8 160000 --- a/3rdparty/QoLA +++ b/3rdparty/QoLA @@ -1 +1 @@ -Subproject commit ed52843ab193a1d8fe63a6a4ebf2198fb0af934d +Subproject commit d48c1d4ed89434e8ccf9e694373ed1d8a187464f diff --git a/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_bwd.cpp b/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_bwd.cpp index cb8b07eb23..128d1d0e52 100644 --- a/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_bwd.cpp +++ b/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_bwd.cpp @@ -531,8 +531,8 @@ BwdFmhaArgs build_bwd_fmha_args(const CkAttnBwdArgs& args){ } aiter::mha_bwd_args fmha_args{}; - fmha_args.sink_ptr = nullptr; - fmha_args.d_sink_ptr = nullptr; + fmha_args.sink_ptr = args.sink_ptr; + fmha_args.d_sink_ptr = args.d_sink_ptr; fmha_args.mask_type = static_cast(static_cast(args.attn_mask_type)); // Mirrors AITER's small-seqlen guard at aiter/ops/mha.py:1689. fmha_args.use_asm_v3 = (args.s_q < 16) ? false : args.uses_bwd_v3; @@ -566,8 +566,6 @@ BwdFmhaArgs build_bwd_fmha_args(const CkAttnBwdArgs& args){ fmha_args.dbias_ptr = ((!args.is_group_mode()) && has_dbias) ? (bias_shape==BiasShape::kBHSS ? args.dbias_ptr : args.dbias_expanded_ptr) : nullptr; - fmha_args.sink_ptr = args.sink_ptr; - fmha_args.d_sink_ptr = args.d_sink_ptr; if (args.is_group_mode()) { fmha_args.seqstart_q_ptr = args.cu_seqlen_q_padded_ptr==nullptr? args.cu_seqlen_q_ptr : args.cu_seqlen_q_padded_ptr; From 85bb3d3940bce36e9415817da0e2039627307bd7 Mon Sep 17 00:00:00 2001 From: shurale-nkn Date: Mon, 27 Jul 2026 13:53:40 +0000 Subject: [PATCH 5/9] bug in TE --- transformer_engine/jax/csrc/extensions/attention.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/transformer_engine/jax/csrc/extensions/attention.cpp b/transformer_engine/jax/csrc/extensions/attention.cpp index cd245e2d07..8062fbdf03 100644 --- a/transformer_engine/jax/csrc/extensions/attention.cpp +++ b/transformer_engine/jax/csrc/extensions/attention.cpp @@ -297,7 +297,7 @@ 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, rng_state, bias, softmax_offset); /* Call the underlying NVTE API */ auto dummy_page_table_tensor = TensorWrapper(nullptr, std::vector{1}, DType::kInt32); From 342335f7c48b28d17fccd3323575f6368c55c816 Mon Sep 17 00:00:00 2001 From: shurale-nkn Date: Wed, 29 Jul 2026 15:10:45 +0000 Subject: [PATCH 6/9] fix for swa + sink --- 3rdparty/QoLA | 2 +- tests/jax/test_fused_attn.py | 4 +++- .../common/fused_attn_rocm/fused_attn_ck.cpp | 23 ++++++++----------- 3 files changed, 14 insertions(+), 15 deletions(-) diff --git a/3rdparty/QoLA b/3rdparty/QoLA index d48c1d4ed8..764da7ac4e 160000 --- a/3rdparty/QoLA +++ b/3rdparty/QoLA @@ -1 +1 @@ -Subproject commit d48c1d4ed89434e8ccf9e694373ed1d8a187464f +Subproject commit 764da7ac4e684b2e78547546407193379c5086a1 diff --git a/tests/jax/test_fused_attn.py b/tests/jax/test_fused_attn.py index 1c34d262ed..61c5f7b30f 100644 --- a/tests/jax/test_fused_attn.py +++ b/tests/jax/test_fused_attn.py @@ -1159,10 +1159,12 @@ def check_dqkv(primitive, reference, pad, idx): jnp.isinf(primitive_dsoftmax_offset) ), "Fused dsoftmax_offset contains Inf" + # softmax_offset is always fp32, but its gradient is only as accurate as the + # attention math that produced it, so tolerance follows the compute dtype. assert_allclose( primitive_dsoftmax_offset, reference_dsoftmax_offset, - dtype=self.softmax_offset.dtype, + dtype=self.dtype, ) if self.coll_count_ref is not None: diff --git a/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp b/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp index 42735f3a1b..47a40c7589 100644 --- a/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp +++ b/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp @@ -77,16 +77,6 @@ bool is_ck_backend_supported( return false; } - // WA for failed validation in CK FMHA with SWA and non-vanilla sink - const bool is_swa = window_size_left >= 0; - if (is_swa && softmax_type != NVTE_VANILLA_SOFTMAX) { - if (nvte_log_ck_config) { - std::cout << "CK fused attn does not support non-vanilla sink with SWA" - << std::endl; - } - return false; - } - // joint filter based on sliding window and attn_mask bool is_causal = is_causal_mask(attn_mask_type); if(is_causal){ @@ -653,8 +643,13 @@ void fused_attn_ck_fwd_impl( const bool has_sink = (softmax_type != NVTE_VANILLA_SOFTMAX && devPtrSoftmaxOffset != nullptr); ck_args.has_sink = has_sink; ck_args.sink_ptr = has_sink ? devPtrSoftmaxOffset : nullptr; - ck_args.sink_size = has_sink ? 1 : 0; - + // sink_size is CK's StreamingLLM sink *prefix width* in key columns, which is a + // different feature from the learnable softmax offset NVTE asks for here. The + // offset only needs has_sink + sink_ptr (CK folds it into the softmax + // denominator). + //(aiter's mha_bwd_args has no sink_size at all). Keep it at 0. + ck_args.sink_size = 0; + // ASM v3 does not support sink; force CK tile path ck_args.uses_fwd_v3 = nvte_ck_uses_fwd_v3 && !has_sink; @@ -1048,7 +1043,9 @@ void fused_attn_ck_bwd_impl( ck_args.how_v3_bf16_cvt = nvte_ck_how_v3_bf16_cvt; ck_args.has_sink = has_sink; ck_args.sink_ptr = has_sink ? devPtrSoftmaxOffset : nullptr; - ck_args.sink_size = has_sink ? 1 : 0; + // See the forward pass: sink_size is the StreamingLLM key-column prefix, not the + // learnable softmax offset, and must stay 0 to keep the mask identical to fwd. + ck_args.sink_size = 0; ck_args.d_sink_ptr = has_sink ? devPtrDSoftmaxOffset : nullptr; if(is_SBHD && is_padding){ From 1120c7649689a494acbe6a22a76da148c75ee721 Mon Sep 17 00:00:00 2001 From: shurale-nkn Date: Wed, 29 Jul 2026 16:39:55 +0000 Subject: [PATCH 7/9] removed fixed code --- .../include/ck_fused_attn/ck_fused_attn.hpp | 1 - .../ck_fused_attn/src/ck_fused_attn_fwd.cpp | 8 +++-- .../common/fused_attn_rocm/fused_attn_ck.cpp | 29 +++---------------- .../jax/csrc/extensions/attention.cpp | 2 +- 4 files changed, 11 insertions(+), 29 deletions(-) diff --git a/transformer_engine/common/ck_fused_attn/include/ck_fused_attn/ck_fused_attn.hpp b/transformer_engine/common/ck_fused_attn/include/ck_fused_attn/ck_fused_attn.hpp index 4ed059b4b5..c5379c41da 100644 --- a/transformer_engine/common/ck_fused_attn/include/ck_fused_attn/ck_fused_attn.hpp +++ b/transformer_engine/common/ck_fused_attn/include/ck_fused_attn/ck_fused_attn.hpp @@ -82,7 +82,6 @@ struct CKAttnCommonArgs { // Softmax sink (learnable / off-by-one) const void* sink_ptr = nullptr; bool has_sink = false; - int64_t sink_size = 0; // O layout (o_ptr lives in derived because fwd writes it / bwd reads it) uint64_t stride_b_o = 0, stride_h_o = 0, stride_s_o = 0; diff --git a/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp b/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp index bdaeaed7ad..9feebb339c 100644 --- a/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp +++ b/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp @@ -92,7 +92,6 @@ void log_fwd_config(const char* func_name, bool has_dropout, const aiter::mha_fw log_value(log_file, "dropout_offset_ptr", std::get<1>(std::get>(fmha_args.drop_seed_offset))); log_value(log_file, "has_sink", fmha_args.has_sink); log_value(log_file, "sink_ptr", fmha_args.sink_ptr); - log_value(log_file, "sink_size", fmha_args.sink_size); } void dump_fwd_timings(const char* dump_path, float average_runtime){ @@ -207,7 +206,12 @@ aiter::mha_fwd_args build_fwd_fmha_args(const CKAttnFwdArgs& args){ fmha_args.q_descale_ptr = nullptr; fmha_args.k_descale_ptr = nullptr; fmha_args.v_descale_ptr = nullptr; - fmha_args.sink_size = args.sink_size; + // sink_size is CK's StreamingLLM sink *prefix width* in key columns, which is a + // different feature from the learnable softmax offset NVTE asks for here. The + // offset only needs has_sink + sink_ptr (CK folds it into the softmax + // denominator). + //(aiter's mha_bwd_args has no sink_size at all). Keep it at 0. + fmha_args.sink_size = 0; fmha_args.min_seqlen_q = 0; fmha_args.block_scale_size_q = 0; fmha_args.block_scale_size_kv = 0; diff --git a/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp b/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp index 47a40c7589..8462997b66 100644 --- a/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp +++ b/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp @@ -137,16 +137,6 @@ bool is_ck_backend_supported( } return false; } - - // WA failed fwd test - bool is_no_mask_window = (window_size_left == -1 && window_size_right == -1); - if(is_ragged && !is_causal && is_no_mask_window && (max_seqlen_q != max_seqlen_kv)){ - if(nvte_log_ck_config){ - std::cout<<"Ragged (THD) cross-attention with plain no-mask padding and " - "max_seqlen_q != max_seqlen_kv is not supported by CK yet"< Date: Thu, 30 Jul 2026 14:54:50 +0000 Subject: [PATCH 8/9] format + small relocation --- tests/jax/test_fused_attn.py | 2 +- .../ck_fused_attn/src/ck_fused_attn_fwd.cpp | 4 +-- .../common/fused_attn_rocm/fused_attn.cpp | 1 - .../common/fused_attn_rocm/fused_attn_ck.cpp | 25 +++++++-------- .../common/fused_attn_rocm/fused_attn_ck.h | 1 - .../jax/csrc/extensions/attention.cpp | 31 ++++++++++--------- 6 files changed, 32 insertions(+), 32 deletions(-) diff --git a/tests/jax/test_fused_attn.py b/tests/jax/test_fused_attn.py index 42b4bf42a1..3a7542d950 100644 --- a/tests/jax/test_fused_attn.py +++ b/tests/jax/test_fused_attn.py @@ -1146,7 +1146,7 @@ def grad_func( grad_shardings = (self.qkvo_sharding, self.qkvo_sharding, self.qkvo_sharding) optional_dgrad_idx = 3 - + # We can compute dBias only for the [1, h, s, s] layout compute_dbias = self.bias_shape == BiasShape._1HSS if compute_dbias: diff --git a/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp b/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp index 9feebb339c..3d60b1695e 100644 --- a/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp +++ b/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp @@ -209,8 +209,8 @@ aiter::mha_fwd_args build_fwd_fmha_args(const CKAttnFwdArgs& args){ // sink_size is CK's StreamingLLM sink *prefix width* in key columns, which is a // different feature from the learnable softmax offset NVTE asks for here. The // offset only needs has_sink + sink_ptr (CK folds it into the softmax - // denominator). - //(aiter's mha_bwd_args has no sink_size at all). Keep it at 0. + // denominator). + // (aiter's mha_bwd_args has no sink_size at all). Keep it at 0. fmha_args.sink_size = 0; fmha_args.min_seqlen_q = 0; fmha_args.block_scale_size_q = 0; diff --git a/transformer_engine/common/fused_attn_rocm/fused_attn.cpp b/transformer_engine/common/fused_attn_rocm/fused_attn.cpp index 64b75905cf..037dc2baa3 100644 --- a/transformer_engine/common/fused_attn_rocm/fused_attn.cpp +++ b/transformer_engine/common/fused_attn_rocm/fused_attn.cpp @@ -309,7 +309,6 @@ NVTE_Fused_Attn_Backend nvte_get_fused_attn_backend( qkv_layout, bias_type, attn_mask_type, - softmax_type, dropout, num_attn_heads, num_gqa_groups, max_seqlen_q, max_seqlen_kv, diff --git a/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp b/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp index 7f4fc2202b..ae86e54194 100644 --- a/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp +++ b/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp @@ -26,7 +26,6 @@ bool is_ck_backend_supported( NVTE_QKV_Layout qkv_layout, NVTE_Bias_Type bias_type, NVTE_Mask_Type attn_mask_type, - NVTE_Softmax_Type softmax_type, float dropout, size_t num_attn_heads, size_t num_gqa_groups, size_t max_seqlen_q, size_t max_seqlen_kv, @@ -468,9 +467,11 @@ void fused_attn_ck_fwd_impl( size_t *workspace_size, cudaStream_t stream){ + const bool has_sink = (softmax_type != NVTE_VANILLA_SOFTMAX && devPtrSoftmaxOffset != nullptr); const bool nvte_log_ck_config = getenv("NVTE_LOG_CK_CONFIG"); - bool nvte_ck_uses_fwd_v3 = getenv("NVTE_CK_USES_FWD_V3", 1); + // ASM v3 does not support sink; force CK tile path + bool nvte_ck_uses_fwd_v3 = getenv("NVTE_CK_USES_FWD_V3", 1) && !has_sink; int nvte_ck_how_v3_bf16_cvt = getenv("NVTE_CK_HOW_V3_BF16_CVT", 1); bool nvte_ck_zero_out_pad = getenv("NVTE_CK_ZERO_OUT_PAD", 1); NVTE_QKV_Format qkv_format = nvte_get_qkv_format(layout); @@ -480,7 +481,7 @@ void fused_attn_ck_fwd_impl( bool is_padding = is_padding_mask(mask_type); bool bshd_to_thd = is_BSHD && is_padding; - const bool has_sink = (softmax_type != NVTE_VANILLA_SOFTMAX && devPtrSoftmaxOffset != nullptr); + // extract the qkv and o storage bytes to allocate buffer for padding removing // b from cu_seqlen is not the actual storage batch for pad_between_seqs case @@ -638,11 +639,10 @@ void fused_attn_ck_fwd_impl( ck_args.window_size_left = window_size_left; ck_args.window_size_right = window_size_right; ck_args.how_v3_bf16_cvt = nvte_ck_how_v3_bf16_cvt; - + ck_args.has_sink = has_sink; ck_args.sink_ptr = has_sink ? devPtrSoftmaxOffset : nullptr; - // ASM v3 does not support sink; force CK tile path - ck_args.uses_fwd_v3 = nvte_ck_uses_fwd_v3 && !has_sink; + ck_args.uses_fwd_v3 = nvte_ck_uses_fwd_v3; if(is_SBHD && is_padding){ // remove padding for q, k, v @@ -728,17 +728,19 @@ void fused_attn_ck_bwd_impl( size_t *workspace_size, cudaStream_t stream) { + const bool has_sink = (softmax_type != NVTE_VANILLA_SOFTMAX && devPtrSoftmaxOffset != nullptr); const bool nvte_log_ck_config = getenv("NVTE_LOG_CK_CONFIG"); // bwd v3 is optional by enabling the following envs // default values follows the ck example setting - bool nvte_ck_uses_bwd_v3 = getenv("NVTE_CK_USES_BWD_V3", 1); + // ASM v3 bwd does not compute the sink gradient; force the CK tile path so + // has_sink users still get a correct (if slower) d_sink, mirroring the fwd + // ASM v3 sink guard above. + bool nvte_ck_uses_bwd_v3 = getenv("NVTE_CK_USES_BWD_V3", 1) && !has_sink; bool nvte_ck_is_v3_atomic_fp32 = getenv("NVTE_CK_IS_V3_ATOMIC_FP32", 1); int nvte_ck_how_v3_bf16_cvt = getenv("NVTE_CK_HOW_V3_BF16_CVT", 1); bool is_mqa_gqa = (h > hg); - const bool has_sink = (softmax_type != NVTE_VANILLA_SOFTMAX && devPtrSoftmaxOffset != nullptr); - NVTE_QKV_Format qkv_format = nvte_get_qkv_format(layout); bool is_ragged = qkv_format==NVTE_QKV_Format::NVTE_THD; bool is_SBHD = qkv_format==NVTE_QKV_Format::NVTE_SBHD || qkv_format==NVTE_QKV_Format::NVTE_SBHD_2BSHD; @@ -1026,10 +1028,7 @@ void fused_attn_ck_bwd_impl( ck_args.aiter_workspace_ptr = aiter_workspace; ck_args.aiter_workspace_bytes = aiter_workspace_bytes; ck_args.deterministic = deterministic; - // ASM v3 bwd does not compute the sink gradient; force the CK tile path so - // has_sink users still get a correct (if slower) d_sink, mirroring the fwd - // ASM v3 sink guard above. - ck_args.uses_bwd_v3 = nvte_ck_uses_bwd_v3 && !has_sink; + ck_args.uses_bwd_v3 = nvte_ck_uses_bwd_v3; ck_args.is_v3_atomic_fp32 = nvte_ck_is_v3_atomic_fp32; ck_args.how_v3_bf16_cvt = nvte_ck_how_v3_bf16_cvt; ck_args.has_sink = has_sink; diff --git a/transformer_engine/common/fused_attn_rocm/fused_attn_ck.h b/transformer_engine/common/fused_attn_rocm/fused_attn_ck.h index 25e91f944f..4e244bb59d 100644 --- a/transformer_engine/common/fused_attn_rocm/fused_attn_ck.h +++ b/transformer_engine/common/fused_attn_rocm/fused_attn_ck.h @@ -22,7 +22,6 @@ bool is_ck_backend_supported( NVTE_QKV_Layout qkv_layout, NVTE_Bias_Type bias_type, NVTE_Mask_Type attn_mask_type, - NVTE_Softmax_Type softmax_type, float dropout, size_t num_attn_heads, size_t num_gqa_groups, size_t max_seqlen_q, size_t max_seqlen_kv, diff --git a/transformer_engine/jax/csrc/extensions/attention.cpp b/transformer_engine/jax/csrc/extensions/attention.cpp index 6fbe750d6e..a6e2e7118e 100644 --- a/transformer_engine/jax/csrc/extensions/attention.cpp +++ b/transformer_engine/jax/csrc/extensions/attention.cpp @@ -51,9 +51,9 @@ void PrepareFusedAttnForwardAuxTensors(NVTETensorPack *tensor_pack, const size_t const size_t bias_batch, const size_t attn_heads, 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, + NVTE_Bias_Type bias_type, NVTE_Softmax_Type softmax_type, + NVTE_Fused_Attn_Backend backend, void *softmax_buf, + 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; @@ -104,8 +104,8 @@ void PrepareFusedAttnForwardAuxTensors(NVTETensorPack *tensor_pack, const size_t bias_aux_data.dtype = static_cast(dtype); nvte_set_tensor_param(&bias_aux, kNVTERowwiseData, &bias_aux_data); } - // include softmax_offset if provided - if (softmax_offset_buf != nullptr) { + // include softmax_offset if the softmax variant carries one + if (softmax_type != NVTE_Softmax_Type::NVTE_VANILLA_SOFTMAX) { NVTETensor &softmax_offset_aux = tensor_pack->tensors[size]; size++; NVTEBasicTensor softmax_offset_aux_data; @@ -135,9 +135,10 @@ void PrepareFusedAttnForwardAuxTensors(NVTETensorPack *tensor_pack, const size_t void PrepareFusedAttnBackwardAuxTensors(NVTETensorPack *tensor_pack, const size_t input_batch, const size_t bias_batch, const size_t attn_heads, const size_t bias_heads, const size_t q_max_seqlen, - const size_t kv_max_seqlen, DType dtype, - [[maybe_unused]]NVTE_Bias_Type bias_type, - [[maybe_unused]]NVTE_Fused_Attn_Backend backend, + const size_t kv_max_seqlen, DType dtype, + [[maybe_unused]] NVTE_Bias_Type bias_type, + NVTE_Softmax_Type softmax_type, + [[maybe_unused]] NVTE_Fused_Attn_Backend backend, void *softmax_buf, void *rng_state_buf, void *bias_buf, void *softmax_offset_buf = nullptr) { @@ -148,11 +149,11 @@ 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, - softmax_offset_buf); + softmax_type, dummy_backend, softmax_buf, rng_state_buf, + bias_buf, softmax_offset_buf); #else PrepareFusedAttnForwardAuxTensors(tensor_pack, input_batch, bias_batch, attn_heads, bias_heads, - q_max_seqlen, kv_max_seqlen, dtype, bias_type, + q_max_seqlen, kv_max_seqlen, dtype, bias_type, softmax_type, backend, softmax_buf, rng_state_buf, bias_buf, softmax_offset_buf); #endif @@ -307,7 +308,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, rng_state, bias, softmax_offset); + softmax_type, backend, softmax_aux, rng_state, bias, + softmax_offset); /* Call the underlying NVTE API */ auto dummy_page_table_tensor = TensorWrapper(nullptr, std::vector{1}, DType::kInt32); @@ -592,8 +594,9 @@ static void FusedAttnBackwardImpl( q_max_seqlen, kv_max_seqlen, qk_head_dim, v_head_dim, window_size_left, window_size_right, false, false, deterministic); PrepareFusedAttnBackwardAuxTensors(&aux_input_tensors, input_batch, bias_batch, attn_heads, - bias_heads, q_max_seqlen, kv_max_seqlen, dtype, bias_type, - backend, softmax_aux, rng_state, bias, softmax_offset); + bias_heads, q_max_seqlen, kv_max_seqlen, dtype, bias_type, + softmax_type, backend, softmax_aux, rng_state, bias, + softmax_offset); /* Call the underlying NVTE API */ // Prepare Q, K, V pointers and shapes based on layout From 17cec0911df89fd4a1612c0d87d681aff5ea4cfd Mon Sep 17 00:00:00 2001 From: shurale-nkn Date: Fri, 31 Jul 2026 16:03:16 +0000 Subject: [PATCH 9/9] CI fix. Fixed hlo check for Distributed test --- tests/jax/test_distributed_fused_attn.py | 39 +++++++++++++++---- .../common/fused_attn_rocm/fused_attn_ck.cpp | 4 +- 2 files changed, 34 insertions(+), 9 deletions(-) diff --git a/tests/jax/test_distributed_fused_attn.py b/tests/jax/test_distributed_fused_attn.py index e1b965a54e..2d500e25b5 100644 --- a/tests/jax/test_distributed_fused_attn.py +++ b/tests/jax/test_distributed_fused_attn.py @@ -52,7 +52,7 @@ class TestDistributedSelfAttn: def generate_collectives_count_ref( - self, mesh_shape, mesh_axes, mesh_resource, with_bias, shape, dtype + self, mesh_shape, mesh_axes, mesh_resource, with_bias, shape, dtype, softmax_type ): jax_dtype = jax.dtypes.canonicalize_dtype(dtype) _, seqlen, heads, _ = shape @@ -64,8 +64,13 @@ def generate_collectives_count_ref( all_reduce_loss_bytes = 4 # 1 * FP32 bias_bytes = int(with_bias) * (heads // tp_size) * seqlen * seqlen * jax_dtype.itemsize - allreduce_total_bytes = all_reduce_loss_bytes + (bias_bytes * is_dp_enabled) - # for loss and dbias + # dsoftmax_offset is [1, heads, 1, 1] and always FP32, regardless of the QKV dtype + with_softmax_offset = softmax_type == AttnSoftmaxType.LEARNABLE_SOFTMAX + softmax_offset_bytes = int(with_softmax_offset) * (heads // tp_size) * 4 + allreduce_total_bytes = all_reduce_loss_bytes + ( + (bias_bytes + softmax_offset_bytes) * is_dp_enabled + ) + # for loss, dbias and dsoftmax_offset return generate_collectives_count(allreduce=allreduce_total_bytes, allgather=0, other=0) def impl_test_self_attn( @@ -111,6 +116,7 @@ def impl_test_self_attn( attn_bias_type != AttnBiasType.NO_BIAS, data_shape, dtype, + softmax_type, ) runner = FusedAttnRunner( batch, @@ -200,10 +206,23 @@ def test_self_attn( class TestDistributedCrossAttn: - def generate_collectives_count_ref(self): - # for loss + def generate_collectives_count_ref( + self, mesh_shape, mesh_axes, mesh_resource, shape, softmax_type + ): + _, _, heads, _ = shape + is_dp_enabled = mesh_resource.dp_resource is not None + tp_size = 1 + if mesh_resource.tpsp_resource is not None: + idx = mesh_axes.index(mesh_resource.tpsp_resource) + tp_size = mesh_shape[idx] + all_reduce_loss_bytes = 4 # 1 * FP32 - return generate_collectives_count(allreduce=all_reduce_loss_bytes, allgather=0, other=0) + # dsoftmax_offset is [1, heads, 1, 1] and always FP32, regardless of the QKV dtype + with_softmax_offset = softmax_type == AttnSoftmaxType.LEARNABLE_SOFTMAX + softmax_offset_bytes = int(with_softmax_offset) * (heads // tp_size) * 4 + allreduce_total_bytes = all_reduce_loss_bytes + (softmax_offset_bytes * is_dp_enabled) + # for loss and dsoftmax_offset + return generate_collectives_count(allreduce=allreduce_total_bytes, allgather=0, other=0) @pytest.mark.parametrize("device_count,mesh_shape,mesh_axes,mesh_resource", generate_configs()) @pytest_parametrize_wrapper("data_shape", DISTRIBUTED_CROSS_ATTN_DATA_SHAPES) @@ -256,7 +275,13 @@ def test_cross_attn( ): pytest.skip("No FusedAttn backend found") - col_ref = self.generate_collectives_count_ref() + col_ref = self.generate_collectives_count_ref( + mesh_shape, + mesh_axes, + mesh_resource, + data_shape, + softmax_type, + ) runner = FusedAttnRunner( batch, seqlen, diff --git a/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp b/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp index ae86e54194..532ba26cfa 100644 --- a/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp +++ b/transformer_engine/common/fused_attn_rocm/fused_attn_ck.cpp @@ -467,7 +467,7 @@ void fused_attn_ck_fwd_impl( size_t *workspace_size, cudaStream_t stream){ - const bool has_sink = (softmax_type != NVTE_VANILLA_SOFTMAX && devPtrSoftmaxOffset != nullptr); + const bool has_sink = softmax_type != NVTE_VANILLA_SOFTMAX; const bool nvte_log_ck_config = getenv("NVTE_LOG_CK_CONFIG"); // ASM v3 does not support sink; force CK tile path @@ -728,7 +728,7 @@ void fused_attn_ck_bwd_impl( size_t *workspace_size, cudaStream_t stream) { - const bool has_sink = (softmax_type != NVTE_VANILLA_SOFTMAX && devPtrSoftmaxOffset != nullptr); + const bool has_sink = softmax_type != NVTE_VANILLA_SOFTMAX; const bool nvte_log_ck_config = getenv("NVTE_LOG_CK_CONFIG"); // bwd v3 is optional by enabling the following envs // default values follows the ck example setting