diff --git a/3rdparty/QoLA b/3rdparty/QoLA index bbcf79e61..764da7ac4 160000 --- a/3rdparty/QoLA +++ b/3rdparty/QoLA @@ -1 +1 @@ -Subproject commit bbcf79e610c69b4d84615c0bb51abfee7bad1a8d +Subproject commit 764da7ac4e684b2e78547546407193379c5086a1 diff --git a/tests/jax/test_distributed_fused_attn.py b/tests/jax/test_distributed_fused_attn.py index e1b965a54..2d500e25b 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/tests/jax/test_fused_attn.py b/tests/jax/test_fused_attn.py index b8c4fb9e2..3a7542d95 100644 --- a/tests/jax/test_fused_attn.py +++ b/tests/jax/test_fused_attn.py @@ -1142,18 +1142,26 @@ def grad_func( } reference_kwargs = {**kwargs, "score_mod_reference": self.score_mod_reference} + 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( @@ -1264,11 +1272,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] @@ -1300,6 +1308,33 @@ 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" + + # 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.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 30d1c1251..c5379c41d 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,10 @@ 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; + // 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; @@ -130,6 +134,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 087a914ab..128d1d0e5 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 @@ -368,6 +368,8 @@ void log_bwd_config(const char* func_name, const aiter::mha_bwd_args& fmha_args, log_value(log_file, "dk_ptr", fmha_args.dk_ptr); log_value(log_file, "dv_ptr", fmha_args.dv_ptr); log_value(log_file, "dbias_ptr", fmha_args.dbias_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); @@ -529,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; 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 59cb92982..3d60b1695 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 @@ -90,6 +90,8 @@ 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); } void dump_fwd_timings(const char* dump_path, float average_runtime){ @@ -153,7 +155,7 @@ aiter::mha_fwd_args build_fwd_fmha_args(const CKAttnFwdArgs& args){ 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; @@ -200,10 +202,15 @@ aiter::mha_fwd_args build_fwd_fmha_args(const CKAttnFwdArgs& args){ 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; + // 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; diff --git a/transformer_engine/common/fused_attn_rocm/fused_attn.cpp b/transformer_engine/common/fused_attn_rocm/fused_attn.cpp index 29b080698..037dc2baa 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, @@ -370,6 +369,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); @@ -401,9 +401,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, @@ -474,6 +474,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; @@ -500,18 +502,26 @@ void nvte_fused_attn_bwd(const NVTETensor Q, const NVTETensor K, const NVTETenso false, 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 7322b96ae..532ba26cf 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, @@ -77,14 +76,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"<("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); @@ -487,6 +482,7 @@ void fused_attn_ck_fwd_impl( bool is_padding = is_padding_mask(mask_type); bool bshd_to_thd = is_BSHD && is_padding; + // 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 size_t q_storage_bytes = max_tokens_q*h*d_qk*nvte_dtype_size(dtype); @@ -625,7 +621,9 @@ void fused_attn_ck_fwd_impl( std::cout<<"(bias_b, bias_h): ("<(dtype, b, h, s_q, d_qk, max_tokens_q, false, q_stride[0], q_stride[1], q_stride[2], devPtrQ, devPtrCuSeqlensQ, devPtrCuSeqlenPaddedQ, devPtrQWithoutPadding, stream); @@ -708,13 +709,16 @@ void fused_attn_ck_bwd_impl( float scaling_factor, float dropout_probability, NVTE_QKV_Layout layout, NVTE_Bias_Type bias_type, NVTE_Mask_Type mask_type, + NVTE_Softmax_Type softmax_type, int64_t window_size_left, int64_t window_size_right, bool deterministic, void* devPtrQ, void* devPtrK, void* devPtrV, void* devPtrO, void* devPtrSoftmaxAux, void* devPtrBias, + void* devPtrSoftmaxOffset, void* devPtrdQ, void* devPtrdK, void* devPtrdV, void* devPtrdO, void* devPtrdBias, + void* devPtrDSoftmaxOffset, void* devPtrDropoutSeed, void* devPtrDropoutOffset, void* devPtrCuSeqlensQ, void* devPtrCuSeqlensKV, @@ -724,10 +728,14 @@ void fused_attn_ck_bwd_impl( size_t *workspace_size, cudaStream_t stream) { + 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 - 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); @@ -923,6 +931,10 @@ void fused_attn_ck_bwd_impl( } // Initialize workspace buffers. + if(has_sink){ + // CK accumulates the sink gradient via atomicAdd (one scalar per head); requires zeroing + NVTE_CHECK_CUDA(cudaMemsetAsync(devPtrDSoftmaxOffset, 0, sizeof(float)*h, stream)); + } // dk_expanded/dv_expanded scratch is written only for valid KV rows by the GQA // main kernel, but the post-dispatch reduction reads every row; pre-zero so // masked/padded rows don't feed uninitialized workspace into dk/dv (-> NaN). @@ -993,7 +1005,10 @@ void fused_attn_ck_bwd_impl( std::cout<<"deterministic: "<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; @@ -1166,53 +1191,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."); } @@ -1223,9 +1246,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), @@ -1247,12 +1271,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, @@ -1279,6 +1306,13 @@ void fused_attn_ck_bwd( bias_h = output_dBias->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; @@ -1302,12 +1336,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 e98281d1a..4e244bb59 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, @@ -36,8 +35,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, @@ -51,12 +52,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 ffb15a63c..a6e2e7118 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,9 +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); } -#ifndef USE_ROCM - // 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; @@ -119,7 +118,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; @@ -138,21 +136,27 @@ void PrepareFusedAttnBackwardAuxTensors(NVTETensorPack *tensor_pack, const size_ 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, + [[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) { +#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_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, softmax_type, + backend, softmax_buf, rng_state_buf, bias_buf, softmax_offset_buf); +#endif } pybind11::tuple GetFusedAttnForwardWorkspaceSizes( @@ -304,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, 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); @@ -589,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, 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