Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 1 addition & 7 deletions onnxruntime/contrib_ops/cpu/bert/attention_common.h
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,6 @@ enum AttentionQkvFormat {
Q_K_V_BNSH, // for non-packed qkv, permuted
Q_K_V_BSNH, // for non-packed qkv, not permuted, used by memory efficient attention or MultiHeadAttention
Q_K_V_BSNH_BNSH_BNSH, // for cross attention, k and v are permuted
Q_K_V_BNSH_QKV_BS3NH, // for TRT fused causal attention, data has two formats (qkv is 3BNSH, gemm_buffer is BS3NH)
Q_K_V_TNH, // for memory efficient attention, qkv are not packed, and paddings are removed.
Q_KV_BSNH_BSN2H, // for TRT fused cross attention, kv are packed
QKV_BSN3H, // for TRT fused attention, qkv are packed
Expand Down Expand Up @@ -106,7 +105,6 @@ enum class AttentionBackend : int {
// The following TRT kernels might be deprecated in the future.
TRT_FLASH_ATTENTION = 32,
TRT_CROSS_ATTENTION = 64,
TRT_CAUSAL_ATTENTION = 128,

// Experimental kernels
LEAN_ATTENTION = 256,
Expand All @@ -122,14 +120,10 @@ constexpr const char* kDisableFusedSelfAttention = "ORT_DISABLE_FUSED_ATTENTION"
// Environment variable to enable or disable fused cross attention kernel. Default is 0 (enabled).
constexpr const char* kDisableFusedCrossAttention = "ORT_DISABLE_FUSED_CROSS_ATTENTION";

// Environment variable to enable or disable TRT fused causal attention kernels. Default is 0 (disabled).
// Note that those causal attention kernels use fp16 accumulation. There is potential accuracy drop using those kernels.
constexpr const char* kEnableFusedCausalAttention = "ORT_ENABLE_FUSED_CAUSAL_ATTENTION";

// Environment variable to enable or disable cuDNN flash attention.
constexpr const char* kEnableCudnnFlashAttention = "ORT_ENABLE_CUDNN_FLASH_ATTENTION";
Comment thread
tianleiwu marked this conversation as resolved.

// Environment variable to enable or disable TRT flash attention. This applies to both self and causal attention. Default is 0 (enabled).
// Environment variable to enable or disable TRT flash attention. Default is 0 (enabled).
constexpr const char* kDisableTrtFlashAttention = "ORT_DISABLE_TRT_FLASH_ATTENTION";

// Environment variable to enable or disable cutlass memory efficient attention. Default is 0 (enabled).
Expand Down
77 changes: 24 additions & 53 deletions onnxruntime/contrib_ops/cuda/bert/attention.cc
Original file line number Diff line number Diff line change
Expand Up @@ -46,10 +46,9 @@ Attention<T>::Attention(const OpKernelInfo& info) : CudaKernel(info), AttentionB
constexpr bool kIsBf16 = std::is_same<T, BFloat16>::value;
constexpr bool kIs16bit = kIsFp16 || kIsBf16;

// We only support FP16 for TRT fused/flash/causal attention.
// We only support FP16 for TRT fused/flash attention.
disable_fused_self_attention_ = !kIsFp16 || !kernel_options_->UseTrtFusedAttention();
enable_trt_flash_attention_ = kIsFp16 && kernel_options_->UseTrtFlashAttention();
enable_fused_causal_attention_ = kIsFp16 && kernel_options_->UseTrtCausalAttention();

disable_memory_efficient_attention_ = kIsBf16 || !kernel_options_->UseEfficientAttention();

Expand Down Expand Up @@ -151,57 +150,29 @@ Status Attention<T>::ComputeInternal(OpKernelContext* context) const {
auto out_accum_buffer = GetScratchBuffer<void>(0, GetComputeStream(context)); // nullptr
#endif

if (!use_flash_attention) {
if (is_unidirectional_) { // GPT
if (enable_fused_causal_attention_) {
// GPT fused kernels requires left side padding. mask can be:
// none (no padding), 1D sequence lengths or 2d mask.
// Fused kernels don't support different sequence lengths of q and kv, so only apply to the first token
// where past state is empty.
bool is_mask_2d_key_padding = parameters.mask_type == AttentionMaskType::MASK_2D_KEY_PADDING;
bool use_causal_fused_runner = (nullptr == mask_index || is_mask_1d_seq_len || is_mask_2d_key_padding) &&
nullptr == attention_bias &&
parameters.past_sequence_length == 0 &&
parameters.hidden_size == parameters.v_hidden_size &&
FusedMHARunnerFP16v2::IsSupported(sm, parameters.head_size, sequence_length,
enable_trt_flash_attention_, true);
if (use_causal_fused_runner) {
// Here we assume that num_heads, head_size and is_unidirectional does not change for an Attention node.
if (nullptr == fused_fp16_runner_.get()) {
std::call_once(fused_fp16_runner_created_, [&]() {
fused_fp16_runner_ = FusedMHARunnerFP16v2::Create(num_heads_, parameters.head_size, sm, is_unidirectional_,
enable_trt_flash_attention_, parameters.scale);
});
}

// Here we assume all causal kernels can be loaded into shared memory. TODO: add a function to check.
fused_runner = fused_fp16_runner_.get();
}
if (!use_flash_attention && !is_unidirectional_) { // BERT
bool use_fused_runner = !disable_fused_self_attention_ &&
(nullptr == mask_index || is_mask_1d_seq_len) &&
nullptr == past &&
nullptr == present &&
nullptr == attention_bias &&
parameters.hidden_size == parameters.v_hidden_size &&
FusedMHARunnerFP16v2::IsSupported(sm, parameters.head_size, sequence_length,
enable_trt_flash_attention_);

if (use_fused_runner) {
// Here we assume that num_heads, head_size and is_unidirectional does not change for an Attention node.
if (nullptr == fused_fp16_runner_.get()) {
std::call_once(fused_fp16_runner_created_, [&]() {
fused_fp16_runner_ = FusedMHARunnerFP16v2::Create(num_heads_, parameters.head_size, sm,
enable_trt_flash_attention_, parameters.scale);
});
}
} else { // BERT
bool use_fused_runner = !disable_fused_self_attention_ &&
(nullptr == mask_index || is_mask_1d_seq_len) &&
nullptr == past &&
nullptr == present &&
nullptr == attention_bias &&
parameters.hidden_size == parameters.v_hidden_size &&
FusedMHARunnerFP16v2::IsSupported(sm, parameters.head_size, sequence_length,
enable_trt_flash_attention_, false);

if (use_fused_runner) {
// Here we assume that num_heads, head_size and is_unidirectional does not change for an Attention node.
if (nullptr == fused_fp16_runner_.get()) {
std::call_once(fused_fp16_runner_created_, [&]() {
fused_fp16_runner_ = FusedMHARunnerFP16v2::Create(num_heads_, parameters.head_size, sm, is_unidirectional_,
enable_trt_flash_attention_, parameters.scale);
});
}

// In case some kernel not loaded due to shared memory limit, we need to double check here.
const int normalized_seq_len = fused_fp16_runner_->NormalizeSequenceLength(sequence_length);
if (fused_fp16_runner_->IsValid(normalized_seq_len)) {
fused_runner = fused_fp16_runner_.get();
}

// In case some kernel not loaded due to shared memory limit, we need to double check here.
const int normalized_seq_len = fused_fp16_runner_->NormalizeSequenceLength(sequence_length);
if (fused_fp16_runner_->IsValid(normalized_seq_len)) {
fused_runner = fused_fp16_runner_.get();
}
}
}
Expand All @@ -227,7 +198,7 @@ Status Attention<T>::ComputeInternal(OpKernelContext* context) const {
debug_info.use_flash_attention = use_flash_attention;
debug_info.use_efficient_attention = use_memory_efficient_attention;
if (fused_runner != nullptr) {
debug_info.SetTrtFusedKernel(is_unidirectional_, enable_trt_flash_attention_, sequence_length);
debug_info.SetTrtFusedKernel(enable_trt_flash_attention_, sequence_length);
}

debug_info.Print("Attention",
Expand Down
1 change: 0 additions & 1 deletion onnxruntime/contrib_ops/cuda/bert/attention.h
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,6 @@ class Attention final : public CudaKernel, public AttentionBase {
bool disable_flash_attention_;
bool disable_fused_self_attention_;
bool enable_trt_flash_attention_;
bool enable_fused_causal_attention_;
bool disable_memory_efficient_attention_;
mutable std::unique_ptr<MHARunner> fused_fp16_runner_;
mutable std::once_flag fused_fp16_runner_created_;
Expand Down
15 changes: 4 additions & 11 deletions onnxruntime/contrib_ops/cuda/bert/attention_impl.cu
Original file line number Diff line number Diff line change
Expand Up @@ -254,7 +254,6 @@ Status FusedTrtSelfAttention(

const int batch_size = parameters.batch_size;
const int sequence_length = parameters.sequence_length;
const bool causal = parameters.is_unidirectional;

const int32_t* sequence_offset = data.cumulated_sequence_length_q_cache;
if (parameters.mask_type == AttentionMaskType::MASK_2D_KEY_PADDING) {
Expand All @@ -274,18 +273,13 @@ Status FusedTrtSelfAttention(

FusedMHARunnerFP16v2* fused_fp16_runner = reinterpret_cast<FusedMHARunnerFP16v2*>(data.fused_runner);

const int s = causal ? sequence_length : fused_fp16_runner->NormalizeSequenceLength(sequence_length);
const int s = fused_fp16_runner->NormalizeSequenceLength(sequence_length);

// B = 2 * batch_size when there is padding in input, and B = batch_size when padding is removed.
const int b = (nullptr == data.mask_index ? batch_size : 2 * batch_size);

if (!causal) {
assert(data.qkv_format == AttentionQkvFormat::QKV_BSN3H);
fused_fp16_runner->Run(b, s, data.q, sequence_offset, data.output, stream);
} else {
assert(data.qkv_format == AttentionQkvFormat::Q_K_V_BNSH_QKV_BS3NH);
fused_fp16_runner->Run(b, s, data.gemm_buffer, sequence_offset, data.output, stream);
}
assert(data.qkv_format == AttentionQkvFormat::QKV_BSN3H);
fused_fp16_runner->Run(b, s, data.q, sequence_offset, data.output, stream);

return Status::OK();
}
Expand Down Expand Up @@ -802,8 +796,7 @@ Status ConcatPastToPresent(int batch_size, int num_heads, int qk_head_size, int
// When there is past state, the head size for Q/K/V shall be same: H == H_v.

if (nullptr != data.present) { // Attention op
assert(data.qkv_format == AttentionQkvFormat::Q_K_V_BNSH ||
data.qkv_format == AttentionQkvFormat::Q_K_V_BNSH_QKV_BS3NH);
assert(data.qkv_format == AttentionQkvFormat::Q_K_V_BNSH);

ORT_RETURN_IF_ERROR(
LaunchConcatTensorToTensor(
Expand Down
11 changes: 2 additions & 9 deletions onnxruntime/contrib_ops/cuda/bert/attention_kernel_options.cc
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,6 @@ void AttentionKernelOptions::Initialize(int value, bool use_build_flag, bool che
use_unfused_ = (value & static_cast<int>(AttentionBackend::MATH)) > 0;
use_trt_flash_attention_ = (value & static_cast<int>(AttentionBackend::TRT_FLASH_ATTENTION)) > 0;
use_trt_cross_attention_ = (value & static_cast<int>(AttentionBackend::TRT_CROSS_ATTENTION)) > 0;
use_trt_causal_attention_ = (value & static_cast<int>(AttentionBackend::TRT_CAUSAL_ATTENTION)) > 0;

use_decoder_attention_ = (value & static_cast<int>(AttentionBackend::DECODER_ATTENTION)) > 0;
} else {
Expand All @@ -50,7 +49,6 @@ void AttentionKernelOptions::Initialize(int value, bool use_build_flag, bool che
use_unfused_ = true;
use_trt_flash_attention_ = !ParseEnvironmentVariableWithDefault<bool>(kDisableTrtFlashAttention, false);
use_trt_cross_attention_ = !ParseEnvironmentVariableWithDefault<bool>(kDisableFusedCrossAttention, false);
use_trt_causal_attention_ = ParseEnvironmentVariableWithDefault<bool>(kEnableFusedCausalAttention, false);

use_decoder_attention_ = !ParseEnvironmentVariableWithDefault<bool>(kDisableDecoderAttention, false);
}
Expand Down Expand Up @@ -112,7 +110,6 @@ void AttentionKernelOptions::Print() const {
sstream << " CUDNN_FLASH_ATTENTION=" << int(use_cudnn_flash_attention_);
sstream << " TRT_FLASH_ATTENTION=" << int(use_trt_flash_attention_);
sstream << " TRT_CROSS_ATTENTION=" << int(use_trt_cross_attention_);
sstream << " TRT_CAUSAL_ATTENTION=" << int(use_trt_causal_attention_);
sstream << " DECODER_ATTENTION=" << int(use_decoder_attention_);
sstream << " MATH=" << int(use_unfused_);

Expand All @@ -126,10 +123,8 @@ void AttentionKernelOptions::Print() const {
}

// Classify the kernel used in TRT fused runner.
void AttentionKernelDebugInfo::SetTrtFusedKernel(bool causal, bool enable_trt_flash_attention, int sequence_length) {
if (causal) {
use_trt_causal_attention = true;
} else if (enable_trt_flash_attention && sequence_length >= contrib::cuda::kMinSequenceLengthFlashAttention) {
void AttentionKernelDebugInfo::SetTrtFusedKernel(bool enable_trt_flash_attention, int sequence_length) {
if (enable_trt_flash_attention && sequence_length >= contrib::cuda::kMinSequenceLengthFlashAttention) {
use_trt_flash_attention = true;
} else {
use_trt_fused_attention = true;
Expand Down Expand Up @@ -172,8 +167,6 @@ void AttentionKernelDebugInfo::Print(const char* operator_name,
sstream << "TRT_FLASH_ATTENTION";
} else if (use_trt_cross_attention.has_value() && use_trt_cross_attention.value()) {
sstream << "TRT_CROSS_ATTENTION";
} else if (use_trt_causal_attention.has_value() && use_trt_causal_attention.value()) {
sstream << "TRT_CAUSAL_ATTENTION";
} else if (use_decoder_attention.has_value() && use_decoder_attention.value()) {
sstream << "DECODER_ATTENTION";
} else {
Expand Down
6 changes: 1 addition & 5 deletions onnxruntime/contrib_ops/cuda/bert/attention_kernel_options.h
Original file line number Diff line number Diff line change
Expand Up @@ -15,9 +15,8 @@ struct AttentionKernelDebugInfo {
std::optional<bool> use_cudnn_flash_attention = std::nullopt;
std::optional<bool> use_trt_flash_attention = std::nullopt;
std::optional<bool> use_trt_cross_attention = std::nullopt;
std::optional<bool> use_trt_causal_attention = std::nullopt;
std::optional<bool> use_decoder_attention = std::nullopt;
void SetTrtFusedKernel(bool causal, bool enable_trt_flash_attention, int sequence_length);
void SetTrtFusedKernel(bool enable_trt_flash_attention, int sequence_length);
void Print(const char* operator_name, const std::string& node_name, bool is_float16, bool is_bfloat16) const;
};

Expand All @@ -33,7 +32,6 @@ class AttentionKernelOptions {
bool UseUnfusedAttention() const { return use_unfused_; }
bool UseTrtFlashAttention() const { return use_trt_flash_attention_; }
bool UseTrtCrossAttention() const { return use_trt_cross_attention_; }
bool UseTrtCausalAttention() const { return use_trt_causal_attention_; }
bool UseDecoderAttention() const { return use_decoder_attention_; }

// True when the SDPA kernel was explicitly selected via the sdpa_kernel provider option
Expand Down Expand Up @@ -67,8 +65,6 @@ class AttentionKernelOptions {

bool use_trt_flash_attention_{true};
bool use_trt_cross_attention_{true};
// Causal attention is disabled by default in #14732.
bool use_trt_causal_attention_{false};

bool use_decoder_attention_{true};

Expand Down
10 changes: 2 additions & 8 deletions onnxruntime/contrib_ops/cuda/bert/attention_prepare_qkv.cu
Original file line number Diff line number Diff line change
Expand Up @@ -125,24 +125,18 @@ Status PrepareQkv_Attention(contrib::AttentionParameters& parameters,
T* qkv = data.workspace;

bool use_fused_kernel = (nullptr != fused_runner && !parameters.is_unidirectional);
bool use_fused_causal = (nullptr != fused_runner && parameters.is_unidirectional);

// For fused TRT attention, transpose qkv to BxSxNx3xH (format 2)
// For flash or memory efficient attention, transpose to 3xBxSxNxH (format 3)
// For unfused kernel, transpose to 3xBxNxSxH (format 1)
// For fused causal kernel, use format 1 since we need have K and V to update present state,
// at the same time, we update gemm_buffer BxSx3xNxH with bias which is used as input for fused causal kernel.
const int format = (use_fused_kernel ? 2 : (use_flash_or_efficient_attention ? 3 : 1));
data.qkv_format = use_fused_kernel
? AttentionQkvFormat::QKV_BSN3H
: (use_flash_or_efficient_attention
? AttentionQkvFormat::Q_K_V_BSNH
: (use_fused_causal
? AttentionQkvFormat::Q_K_V_BNSH_QKV_BS3NH
: AttentionQkvFormat::Q_K_V_BNSH));
: AttentionQkvFormat::Q_K_V_BNSH);

// For fused causal, we will update gemm_buffer with bias directly.
T* qkv_add_bias = use_fused_causal ? data.gemm_buffer : nullptr;
T* qkv_add_bias = nullptr;

int matrix_to_transpose = ((format == AttentionQkvFormat::Q_K_V_BNSH && past_present_share_buffer) ? 1 : 3);
// format 1: BxSx(NH + NH + NH_v) => BxNxSxH + BxNxSxH + BxNxSxH_v
Expand Down
6 changes: 3 additions & 3 deletions onnxruntime/contrib_ops/cuda/bert/multihead_attention.cc
Original file line number Diff line number Diff line change
Expand Up @@ -432,14 +432,14 @@ Status MultiHeadAttention<T, QK>::ComputeInternal(OpKernelContext* context) cons
parameters.hidden_size == parameters.v_hidden_size &&
parameters.sequence_length == parameters.kv_sequence_length && // self attention only for fused runner
FusedMHARunnerFP16v2::IsSupported(sm, parameters.head_size, sequence_length,
enable_trt_flash_attention_, is_unidirectional_);
enable_trt_flash_attention_);

DUMP_STRING("Use fused runner = ", (use_fused_runner == true));
if (use_fused_runner) {
// Here we assume that num_heads and head_size does not change for a MultiHeadAttention node.
if (nullptr == fused_fp16_runner_.get()) {
std::call_once(fused_fp16_runner_created_, [&]() {
fused_fp16_runner_ = FusedMHARunnerFP16v2::Create(num_heads_, parameters.head_size, sm, is_unidirectional_,
fused_fp16_runner_ = FusedMHARunnerFP16v2::Create(num_heads_, parameters.head_size, sm,
enable_trt_flash_attention_, parameters.scale);
});
}
Expand Down Expand Up @@ -568,7 +568,7 @@ Status MultiHeadAttention<T, QK>::ComputeInternal(OpKernelContext* context) cons
debug_info.use_trt_cross_attention = fused_cross_attention_kernel != nullptr;
debug_info.use_efficient_attention = use_memory_efficient_attention;
if (fused_fp16_runner_ != nullptr) {
debug_info.SetTrtFusedKernel(is_unidirectional_, enable_trt_flash_attention_, sequence_length);
debug_info.SetTrtFusedKernel(enable_trt_flash_attention_, sequence_length);
}
debug_info.Print("MultiHeadAttention",
this->Node().Name(),
Expand Down
7 changes: 3 additions & 4 deletions onnxruntime/contrib_ops/cuda/bert/packed_attention.cc
Original file line number Diff line number Diff line change
Expand Up @@ -60,16 +60,15 @@ MHARunner* TrtFusedAttention<T>::GetFusedRunner(const cudaDeviceProp& device_pro
bool is_fMHA_supported = FusedMHARunnerFP16v2::IsSupported(sm,
parameters.head_size,
parameters.sequence_length,
enable_trt_flash_attention_,
false /*causal*/);
enable_trt_flash_attention_);

if (!is_fMHA_supported) {
return fused_runner;
}

// Assuming that num_heads and head_size do not change.
if (nullptr == fused_fp16_runner_.get()) {
fused_fp16_runner_ = FusedMHARunnerFP16v2::Create(parameters.num_heads, parameters.head_size, sm, false /*causal*/,
fused_fp16_runner_ = FusedMHARunnerFP16v2::Create(parameters.num_heads, parameters.head_size, sm,
enable_trt_flash_attention_, parameters.scale);
}

Expand Down Expand Up @@ -269,7 +268,7 @@ Status PackedAttention<T>::ComputeInternal(OpKernelContext* context) const {
AttentionKernelDebugInfo debug_info;
debug_info.use_efficient_attention = use_memory_efficient_attention;
if (fused_runner != nullptr) {
debug_info.SetTrtFusedKernel(false /*causal*/, this->enable_trt_flash_attention_, parameters.sequence_length);
debug_info.SetTrtFusedKernel(this->enable_trt_flash_attention_, parameters.sequence_length);
}

debug_info.Print("PackedAttention",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -240,7 +240,7 @@ Status PackedMultiHeadAttention<T>::ComputeInternal(OpKernelContext* context) co
debug_info.use_flash_attention = use_flash_attention;
debug_info.use_efficient_attention = use_memory_efficient_attention;
if (fused_runner != nullptr) {
debug_info.SetTrtFusedKernel(false /*causal*/, this->enable_trt_flash_attention_, parameters.sequence_length);
debug_info.SetTrtFusedKernel(this->enable_trt_flash_attention_, parameters.sequence_length);
}

debug_info.Print("PackedMultiHeadAttention",
Expand Down
Loading
Loading