diff --git a/docs/OperatorKernels.md b/docs/OperatorKernels.md index d9efd7a09099e..04ac90d7d231f 100644 --- a/docs/OperatorKernels.md +++ b/docs/OperatorKernels.md @@ -587,6 +587,8 @@ The **OpSet Version** column uses the following notation: |FusedConv|*in* X:**T**
*in* W:**T**
*in* B:**T**
*in* Z:**T**
*out* Y:**T**|1+|**T** = tensor(float)| |FusedGemm|*in* A:**T**
*in* B:**T**
*in* C:**T**
*out* Y:**T**|1+|**T** = tensor(float)| |FusedMatMul|*in* A:**T**
*in* B:**T**
*out* Y:**T**|1+|**T** = tensor(double), tensor(float)| +|GatedAdd|*in* X:**T**
*in* Y:**T**
*in* gate:**T**
*out* output:**T**|1+|**T** = tensor(float), tensor(float16)| +|GatedRMSNorm|*in* X:**T**
*in* scale:**T**
*in* gate:**T**
*out* Y:**T**|1+|**T** = tensor(float), tensor(float16)| |GatherBlockQuantized|*in* data:**T1**
*in* indices:**Tind**
*in* scales:**T2**
*in* zero_points:**T1**
*out* output:**T2**|1+|**T1** = tensor(int4), tensor(uint4), tensor(uint8)
**T2** = tensor(float), tensor(float16)
**Tind** = tensor(int32), tensor(int64)| |GatherND|*in* data:**T**
*in* indices:**Tind**
*out* output:**T**|1+|**T** = tensor(bfloat16), tensor(bool), tensor(double), tensor(float), tensor(float16), tensor(int16), tensor(int32), tensor(int64), tensor(int8), tensor(string), tensor(uint16), tensor(uint32), tensor(uint64), tensor(uint8)
**Tind** = tensor(int32), tensor(int64)| |Gelu|*in* X:**T**
*out* Y:**T**|1+|**T** = tensor(float)| @@ -595,6 +597,7 @@ The **OpSet Version** column uses the following notation: |GroupQueryAttention|*in* query:**T**
*in* key:**T**
*in* value:**T**
*in* past_key:**T_CACHE**
*in* past_value:**T_CACHE**
*in* seqlens_k:**M**
*in* total_sequence_length:**M**
*in* cos_cache:**T**
*in* sin_cache:**T**
*in* position_ids:**tensor(int64)**
*in* attention_bias:**T**
*in* head_sink:**T**
*in* k_scale:**T_KV_SCALE**
*in* v_scale:**T_KV_SCALE**
*in* q_norm_weight:**T**
*in* k_norm_weight:**T**
*out* output:**T**
*out* present_key:**T_CACHE**
*out* present_value:**T_CACHE**
*out* output_qk:**T**|1+|**M** = tensor(int32)
**T** = tensor(float), tensor(float16)
**T_CACHE** = tensor(float), tensor(float16), tensor(int8), tensor(uint8)
**T_KV_SCALE** = tensor(float)| |Inverse|*in* X:**T**
*out* Y:**T**|1+|**T** = tensor(double), tensor(float), tensor(float16)| |LinearAttention|*in* query:**T**
*in* key:**T**
*in* value:**T**
*in* past_state:**S**
*in* decay:**T**
*in* beta:**T**
*out* output:**T**
*out* present_state:**S**|1+|**T** = tensor(float)| +|LinearAttentionGate|*in* a:**T**
*in* dt_bias:**TF**
*in* decay_scale:**TF**
*in* b:**T**
*out* decay:**T**
*out* beta:**T**|1+|**T** = tensor(float), tensor(float16)
**TF** = tensor(float)| |MRotaryEmbedding|*in* input:**T**
*in* position_ids:**M**
*in* cos_cache:**T**
*in* sin_cache:**T**
*out* output:**T**|1+|**M** = tensor(int64)
**T** = tensor(float), tensor(float16)| |MatMulBnb4|*in* A:**T1**
*in* B:**T2**
*in* absmax:**T1**
*out* Y:**T1**|1+|**T1** = tensor(float)
**T2** = tensor(uint8)| |MatMulFpQ4|*in* A:**T1**
*in* B:**T2**
*in* B_shape:**T3**
*out* Y:**T1**|1+|**T1** = tensor(float)
**T2** = tensor(uint8)
**T3** = tensor(int64)| @@ -836,7 +839,7 @@ The **OpSet Version** column uses the following notation: |||[13, 18]|**B** = tensor(bool)
**I** = tensor(int64)
**V** = seq(tensor(bfloat16)), seq(tensor(bool)), seq(tensor(double)), seq(tensor(float)), seq(tensor(float16)), seq(tensor(int16)), seq(tensor(int32)), seq(tensor(int64)), seq(tensor(int8)), seq(tensor(string)), seq(tensor(uint16)), seq(tensor(uint32)), seq(tensor(uint64)), seq(tensor(uint8)), tensor(bfloat16), tensor(bool), tensor(double), tensor(float), tensor(float16), tensor(int16), tensor(int32), tensor(int64), tensor(int8), tensor(string), tensor(uint16), tensor(uint32), tensor(uint64), tensor(uint8)| |||[11, 12]|**B** = tensor(bool)
**I** = tensor(int64)
**V** = tensor(bfloat16), tensor(bool), tensor(double), tensor(float), tensor(float16), tensor(int16), tensor(int32), tensor(int64), tensor(int8), tensor(uint16), tensor(uint32), tensor(uint64), tensor(uint8)| |||[1, 10]|**B** = tensor(bool)
**I** = tensor(int64)
**V** = tensor(bfloat16), tensor(bool), tensor(double), tensor(float), tensor(float16), tensor(int16), tensor(int32), tensor(int64), tensor(int8), tensor(uint16), tensor(uint32), tensor(uint64), tensor(uint8)| -|LpNormalization|*in* input:**T**
*out* output:**T**|22+|**T** = tensor(float), tensor(float16)| +|LpNormalization|*in* input:**T**
*out* output:**T**|22+|**T** = tensor(bfloat16), tensor(float), tensor(float16)| |||[1, 21]|**T** = tensor(float), tensor(float16)| |MatMul|*in* A:**T**
*in* B:**T**
*out* Y:**T**|13+|**T** = tensor(bfloat16), tensor(double), tensor(float), tensor(float16)| |||[9, 12]|**T** = tensor(double), tensor(float), tensor(float16)| @@ -1073,7 +1076,7 @@ The **OpSet Version** column uses the following notation: |BiasSplitGelu|*in* X:**T**
*in* bias:**T**
*out* Y:**T**|1+|**T** = tensor(float), tensor(float16)| |BitmaskBiasDropout|*in* data:**T**
*in* bias:**T**
*in* residual:**T**
*in* ratio:**T1**
*in* training_mode:**T2**
*out* output:**T**
*out* mask:**T3**|1+|**T** = tensor(bfloat16), tensor(double), tensor(float), tensor(float16)
**T1** = tensor(bfloat16), tensor(double), tensor(float), tensor(float16)
**T2** = tensor(bool)
**T3** = tensor(uint32)| |BitmaskDropout|*in* data:**T**
*in* ratio:**T1**
*in* training_mode:**T2**
*out* output:**T**
*out* mask:**T3**|1+|**T** = tensor(bfloat16), tensor(double), tensor(float), tensor(float16)
**T1** = tensor(bfloat16), tensor(double), tensor(float), tensor(float16)
**T2** = tensor(bool)
**T3** = tensor(uint32)| -|CausalConvWithState|*in* input:**T**
*in* weight:**T**
*in* bias:**T**
*in* past_state:**T**
*out* output:**T**
*out* present_state:**T**|1+|**T** = tensor(float), tensor(float16)| +|CausalConvWithState|*in* input:**T**
*in* weight:**T**
*in* bias:**T**
*in* past_state:**T**
*out* output:**T**
*out* present_state:**T**|1+|**T** = tensor(bfloat16), tensor(float), tensor(float16)| |ComplexMul|*in* A:**T**
*in* B:**T**
*out* C:**T**|1+|**T** = tensor(float), tensor(float16)| |ComplexMulConj|*in* A:**T**
*in* B:**T**
*out* C:**T**|1+|**T** = tensor(float), tensor(float16)| |ConvTransposeWithDynamicPads|*in* X:**T**
*in* W:**T**
*in* Pads:**tensor(int64)**
*in* B:**T**
*out* Y:**T**|1+|**T** = tensor(float)| @@ -1100,7 +1103,7 @@ The **OpSet Version** column uses the following notation: |GroupQueryAttention|*in* query:**T**
*in* key:**T**
*in* value:**T**
*in* past_key:**T_CACHE**
*in* past_value:**T_CACHE**
*in* seqlens_k:**M**
*in* total_sequence_length:**M**
*in* cos_cache:**T**
*in* sin_cache:**T**
*in* position_ids:**tensor(int64)**
*in* attention_bias:**T**
*in* head_sink:**T**
*in* k_scale:**T_KV_SCALE**
*in* v_scale:**T_KV_SCALE**
*in* q_norm_weight:**T**
*in* k_norm_weight:**T**
*out* output:**T**
*out* present_key:**T_CACHE**
*out* present_value:**T_CACHE**
*out* output_qk:**T**|1+|**M** = tensor(int32)
**T** = tensor(bfloat16), tensor(float16)
**T_CACHE** = tensor(bfloat16), tensor(float16), tensor(float8e4m3fn), tensor(int8)
**T_KV_SCALE** = tensor(float)| |Inverse|*in* X:**T**
*out* Y:**T**|1+|**T** = tensor(double), tensor(float), tensor(float16)| |Irfft|*in* X:**T**
*out* Y:**T**|1+|**T** = tensor(double), tensor(float), tensor(float16)| -|LinearAttention|*in* query:**T**
*in* key:**T**
*in* value:**T**
*in* past_state:**S**
*in* decay:**T**
*in* beta:**T**
*out* output:**T**
*out* present_state:**S**|1+|**T** = tensor(float), tensor(float16)| +|LinearAttention|*in* query:**T**
*in* key:**T**
*in* value:**T**
*in* past_state:**S**
*in* decay:**T**
*in* beta:**T**
*out* output:**T**
*out* present_state:**S**|1+|**T** = tensor(bfloat16), tensor(float), tensor(float16)| |LinearAttentionGate|*in* a:**T**
*in* dt_bias:**TF**
*in* decay_scale:**TF**
*in* b:**T**
*out* decay:**T**
*out* beta:**T**|1+|**T** = tensor(bfloat16), tensor(float), tensor(float16)
**TF** = tensor(float)| |LongformerAttention|*in* input:**T**
*in* weight:**T**
*in* bias:**T**
*in* mask:**T**
*in* global_weight:**T**
*in* global_bias:**T**
*in* global:**G**
*out* output:**T**|1+|**T** = tensor(float), tensor(float16)| |MRotaryEmbedding|*in* input:**T**
*in* position_ids:**M**
*in* cos_cache:**T**
*in* sin_cache:**T**
*out* output:**T**|1+|**M** = tensor(int64)
**T** = tensor(bfloat16), tensor(float), tensor(float16)| diff --git a/onnxruntime/contrib_ops/cpu/bert/gated_add.cc b/onnxruntime/contrib_ops/cpu/bert/gated_add.cc new file mode 100644 index 0000000000000..48b643da3bc20 --- /dev/null +++ b/onnxruntime/contrib_ops/cpu/bert/gated_add.cc @@ -0,0 +1,95 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "contrib_ops/cpu/bert/gated_add.h" + +#include "core/framework/tensor.h" +#include "core/platform/threadpool.h" + +namespace onnxruntime { +namespace contrib { + +#define REGISTER_KERNEL_TYPED(T) \ + ONNX_OPERATOR_TYPED_KERNEL_EX( \ + GatedAdd, \ + kMSDomain, \ + 1, \ + T, \ + kCpuExecutionProvider, \ + KernelDefBuilder() \ + .TypeConstraint("T", DataTypeImpl::GetTensorType()), \ + GatedAdd); + +REGISTER_KERNEL_TYPED(float) +REGISTER_KERNEL_TYPED(MLFloat16) + +#undef REGISTER_KERNEL_TYPED + +namespace { + +// output = x + round_to_T(y * gate). For MLFloat16 the product is rounded to half before the +// add, matching separate ONNX Mul and Add operators. +template +inline T GatedAddValue(T x, T y, T gate) { + if constexpr (std::is_same_v) { + const T product(y.ToFloat() * gate.ToFloat()); + return T(x.ToFloat() + product.ToFloat()); + } else { + return x + y * gate; + } +} + +} // namespace + +template +Status GatedAdd::Compute(OpKernelContext* context) const { + const Tensor* x = context->Input(0); + const Tensor* y = context->Input(1); + const Tensor* gate = context->Input(2); + const TensorShape& shape = x->Shape(); + + ORT_RETURN_IF_NOT(shape.NumDimensions() >= 1, "X must have rank >= 1"); + ORT_RETURN_IF_NOT(y->Shape() == shape, "Y must have the same shape as X"); + ORT_RETURN_IF_NOT(gate->Shape().NumDimensions() == shape.NumDimensions(), + "gate must have the same rank as X"); + + const size_t last_axis = shape.NumDimensions() - 1; + const int64_t hidden_size = shape[last_axis]; + ORT_RETURN_IF_NOT(hidden_size > 0, "X last dimension must be positive"); + ORT_RETURN_IF_NOT(gate->Shape()[last_axis] == 1, "gate last dimension must be 1"); + for (size_t axis = 0; axis < last_axis; ++axis) { + ORT_RETURN_IF_NOT(gate->Shape()[axis] == shape[axis], + "gate dimension ", axis, " must match X"); + } + + Tensor* output = context->Output(0, shape); + const int64_t count = shape.Size(); + if (count == 0) { + return Status::OK(); + } + + const T* x_data = x->Data(); + const T* y_data = y->Data(); + const T* gate_data = gate->Data(); + T* output_data = output->MutableData(); + const int64_t num_rows = count / hidden_size; + + concurrency::ThreadPool::TryBatchParallelFor( + context->GetOperatorThreadPool(), onnxruntime::narrow(num_rows), + [&](ptrdiff_t row) { + const int64_t offset = row * hidden_size; + const T gate_value = gate_data[row]; + for (int64_t i = 0; i < hidden_size; ++i) { + output_data[offset + i] = GatedAddValue(x_data[offset + i], y_data[offset + i], gate_value); + } + }, + 0); + + return Status::OK(); +} + +template class GatedAdd; +template class GatedAdd; + +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/cpu/bert/gated_add.h b/onnxruntime/contrib_ops/cpu/bert/gated_add.h new file mode 100644 index 0000000000000..b0008e29b3cc1 --- /dev/null +++ b/onnxruntime/contrib_ops/cpu/bert/gated_add.h @@ -0,0 +1,21 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include "core/common/common.h" +#include "core/framework/op_kernel.h" + +namespace onnxruntime { +namespace contrib { + +// output = X + round_to_T(Y * gate), with gate broadcast across the last dimension. +template +class GatedAdd final : public OpKernel { + public: + explicit GatedAdd(const OpKernelInfo& info) : OpKernel(info) {} + Status Compute(OpKernelContext* context) const override; +}; + +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/cpu/bert/linear_attention_gates.cc b/onnxruntime/contrib_ops/cpu/bert/linear_attention_gates.cc new file mode 100644 index 0000000000000..165ad049f0a04 --- /dev/null +++ b/onnxruntime/contrib_ops/cpu/bert/linear_attention_gates.cc @@ -0,0 +1,181 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "contrib_ops/cpu/bert/linear_attention_gates.h" + +#include + +#include "core/framework/tensor.h" +#include "core/mlas/inc/mlas.h" +#include "core/platform/threadpool.h" + +namespace onnxruntime { +namespace contrib { + +#define REGISTER_KERNEL_TYPED(Op, T) \ + ONNX_OPERATOR_TYPED_KERNEL_EX( \ + Op, \ + kMSDomain, \ + 1, \ + T, \ + kCpuExecutionProvider, \ + KernelDefBuilder() \ + .TypeConstraint("T", DataTypeImpl::GetTensorType()) \ + .TypeConstraint("TF", DataTypeImpl::GetTensorType()), \ + Op); + +REGISTER_KERNEL_TYPED(LinearAttentionGate, float) +REGISTER_KERNEL_TYPED(LinearAttentionGate, MLFloat16) + +#undef REGISTER_KERNEL_TYPED + +#define REGISTER_KERNEL_TYPED(Op, T) \ + ONNX_OPERATOR_TYPED_KERNEL_EX( \ + Op, \ + kMSDomain, \ + 1, \ + T, \ + kCpuExecutionProvider, \ + KernelDefBuilder() \ + .TypeConstraint("T", DataTypeImpl::GetTensorType()), \ + Op); + +REGISTER_KERNEL_TYPED(GatedRMSNorm, float) +REGISTER_KERNEL_TYPED(GatedRMSNorm, MLFloat16) + +#undef REGISTER_KERNEL_TYPED + +namespace { + +inline float SigmoidFloat(float value) { + float output; + MlasComputeLogistic(&value, &output, 1); + return output; +} + +inline float SoftplusFloat(float value) { + return value > 0.0f ? value + std::log(std::exp(-value) + 1.0f) : std::log(std::exp(value) + 1.0f); +} + +} // namespace + +template +Status LinearAttentionGate::Compute(OpKernelContext* context) const { + const Tensor* a = context->Input(0); + const Tensor* dt_bias = context->Input(1); + const Tensor* decay_scale = context->Input(2); + const Tensor* b = context->Input(3); // optional + + const auto& a_shape = a->Shape(); + ORT_RETURN_IF_NOT(a_shape.NumDimensions() >= 1, "a must have rank >= 1"); + const int64_t num_heads = a_shape[a_shape.NumDimensions() - 1]; + ORT_RETURN_IF_NOT(num_heads > 0, "a last dimension must be positive"); + + ORT_RETURN_IF_NOT(dt_bias->Shape().Size() == num_heads, + "dt_bias must have ", num_heads, " elements, got ", dt_bias->Shape().Size()); + ORT_RETURN_IF_NOT(decay_scale->Shape().Size() == num_heads, + "decay_scale must have ", num_heads, " elements, got ", decay_scale->Shape().Size()); + + Tensor* decay = context->Output(0, a_shape); + Tensor* beta = context->Output(1, a_shape); + + if (beta != nullptr) { + ORT_RETURN_IF_NOT(b != nullptr, "The b input is required when the beta output is requested"); + ORT_RETURN_IF_NOT(b->Shape() == a_shape, "b must have the same shape as a"); + } + + const int64_t count = a_shape.Size(); + if (count == 0) { + return Status::OK(); + } + + const T* a_data = a->Data(); + const T* b_data = b == nullptr ? nullptr : b->Data(); + const float* dt_bias_data = dt_bias->Data(); + const float* decay_scale_data = decay_scale->Data(); + T* decay_data = decay->MutableData(); + T* beta_data = beta == nullptr ? nullptr : beta->MutableData(); + + const int64_t num_tokens = count / num_heads; + + concurrency::ThreadPool::TryBatchParallelFor( + context->GetOperatorThreadPool(), onnxruntime::narrow(num_tokens), + [&](ptrdiff_t token) { + const int64_t offset = token * num_heads; + for (int64_t h = 0; h < num_heads; ++h) { + const int64_t idx = offset + h; + const float biased = static_cast(a_data[idx]) + dt_bias_data[h]; + decay_data[idx] = static_cast(decay_scale_data[h] * SoftplusFloat(biased)); + if (beta_data != nullptr) { + beta_data[idx] = static_cast(SigmoidFloat(static_cast(b_data[idx]))); + } + } + }, + 0); + + return Status::OK(); +} + +template +GatedRMSNorm::GatedRMSNorm(const OpKernelInfo& info) : OpKernel(info) { + epsilon_ = info.GetAttrOrDefault("epsilon", 1e-5f); +} + +template +Status GatedRMSNorm::Compute(OpKernelContext* context) const { + const Tensor* input = context->Input(0); + const Tensor* scale = context->Input(1); + const Tensor* gate = context->Input(2); + + const auto& shape = input->Shape(); + ORT_RETURN_IF_NOT(shape.NumDimensions() >= 1, "X must have rank >= 1"); + ORT_RETURN_IF_NOT(gate->Shape() == shape, "gate must have the same shape as X"); + + const int64_t norm_size = scale->Shape().Size(); + ORT_RETURN_IF_NOT(norm_size > 0, "scale must not be empty"); + const int64_t last_dim = shape[shape.NumDimensions() - 1]; + ORT_RETURN_IF_NOT(last_dim % norm_size == 0, + "X last dimension (", last_dim, ") must be a multiple of the scale length (", + norm_size, ")"); + + Tensor* output = context->Output(0, shape); + const int64_t count = shape.Size(); + if (count == 0) { + return Status::OK(); + } + const int64_t num_rows = count / norm_size; + + const T* input_data = input->Data(); + const T* scale_data = scale->Data(); + const T* gate_data = gate->Data(); + T* output_data = output->MutableData(); + + concurrency::ThreadPool::TryBatchParallelFor( + context->GetOperatorThreadPool(), onnxruntime::narrow(num_rows), + [&](ptrdiff_t row) { + const int64_t offset = row * norm_size; + float sum_sq = 0.0f; + for (int64_t i = 0; i < norm_size; ++i) { + const float v = static_cast(input_data[offset + i]); + sum_sq += v * v; + } + const float inv_rms = 1.0f / std::sqrt(sum_sq / static_cast(norm_size) + epsilon_); + for (int64_t i = 0; i < norm_size; ++i) { + const float z = static_cast(gate_data[offset + i]); + const float normalized = static_cast(input_data[offset + i]) * inv_rms * + static_cast(scale_data[i]); + output_data[offset + i] = static_cast(normalized * (z * SigmoidFloat(z))); + } + }, + 0); + + return Status::OK(); +} + +template class LinearAttentionGate; +template class LinearAttentionGate; +template class GatedRMSNorm; +template class GatedRMSNorm; + +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/cpu/bert/linear_attention_gates.h b/onnxruntime/contrib_ops/cpu/bert/linear_attention_gates.h new file mode 100644 index 0000000000000..eb3c4b68f31e9 --- /dev/null +++ b/onnxruntime/contrib_ops/cpu/bert/linear_attention_gates.h @@ -0,0 +1,32 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include "core/common/common.h" +#include "core/framework/op_kernel.h" + +namespace onnxruntime { +namespace contrib { + +// decay = decay_scale * Softplus(a + dt_bias), beta = Sigmoid(b). +template +class LinearAttentionGate final : public OpKernel { + public: + explicit LinearAttentionGate(const OpKernelInfo& info) : OpKernel(info) {} + Status Compute(OpKernelContext* context) const override; +}; + +// Y = X * rsqrt(mean(X^2) + epsilon) * scale * SiLU(gate). +template +class GatedRMSNorm final : public OpKernel { + public: + explicit GatedRMSNorm(const OpKernelInfo& info); + Status Compute(OpKernelContext* context) const override; + + private: + float epsilon_; +}; + +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/cpu/cpu_contrib_kernels.cc b/onnxruntime/contrib_ops/cpu/cpu_contrib_kernels.cc index af09ad47e34fa..1d323b18af6fe 100644 --- a/onnxruntime/contrib_ops/cpu/cpu_contrib_kernels.cc +++ b/onnxruntime/contrib_ops/cpu/cpu_contrib_kernels.cc @@ -33,6 +33,12 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, float, SparseAttention); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, MLFloat16, SparseAttention); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, float, LinearAttention); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, float, GatedAdd); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, MLFloat16, GatedAdd); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, float, LinearAttentionGate); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, MLFloat16, LinearAttentionGate); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, float, GatedRMSNorm); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, MLFloat16, GatedRMSNorm); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, float, CausalConvWithState); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, float, RotaryEmbedding); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, MLFloat16, RotaryEmbedding); @@ -331,6 +337,12 @@ Status RegisterCpuContribKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, diff --git a/onnxruntime/contrib_ops/cuda/bert/causal_conv_with_state.cc b/onnxruntime/contrib_ops/cuda/bert/causal_conv_with_state.cc index 8d422e126675e..950fe526aa0f6 100644 --- a/onnxruntime/contrib_ops/cuda/bert/causal_conv_with_state.cc +++ b/onnxruntime/contrib_ops/cuda/bert/causal_conv_with_state.cc @@ -26,6 +26,7 @@ using namespace onnxruntime::cuda; // CudaKernel, Stream, GetDeviceProp, ToCuda REGISTER_KERNEL_TYPED(float) REGISTER_KERNEL_TYPED(MLFloat16) +REGISTER_KERNEL_TYPED(BFloat16) template CausalConvWithState::CausalConvWithState(const OpKernelInfo& info) : CudaKernel(info) { diff --git a/onnxruntime/contrib_ops/cuda/bert/linear_attention.cc b/onnxruntime/contrib_ops/cuda/bert/linear_attention.cc index d34897c3f163d..5111caa9af894 100644 --- a/onnxruntime/contrib_ops/cuda/bert/linear_attention.cc +++ b/onnxruntime/contrib_ops/cuda/bert/linear_attention.cc @@ -29,6 +29,7 @@ using namespace onnxruntime::cuda; // CudaKernel, Stream, GetDeviceProp, ToCuda REGISTER_KERNEL_TYPED(float) REGISTER_KERNEL_TYPED(MLFloat16) +REGISTER_KERNEL_TYPED(BFloat16) template LinearAttention::LinearAttention(const OpKernelInfo& info) : CudaKernel(info) { diff --git a/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc b/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc index e3a517da0bce8..0e092eb36e6a3 100644 --- a/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc +++ b/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc @@ -158,6 +158,7 @@ class CUDA_MS_OP_TYPED_CLASS_NAME(1, BFloat16, MRotaryEmbedding); class CUDA_MS_OP_TYPED_CLASS_NAME(1, MLFloat16, GemmaRotaryEmbedding); class CUDA_MS_OP_TYPED_CLASS_NAME(1, float, LinearAttention); class CUDA_MS_OP_TYPED_CLASS_NAME(1, MLFloat16, LinearAttention); +class CUDA_MS_OP_TYPED_CLASS_NAME(1, BFloat16, LinearAttention); class CUDA_MS_OP_TYPED_CLASS_NAME(1, float, LinearAttentionGate); class CUDA_MS_OP_TYPED_CLASS_NAME(1, MLFloat16, LinearAttentionGate); class CUDA_MS_OP_TYPED_CLASS_NAME(1, BFloat16, LinearAttentionGate); @@ -169,6 +170,7 @@ class CUDA_MS_OP_TYPED_CLASS_NAME(1, MLFloat16, GatedAdd); class CUDA_MS_OP_TYPED_CLASS_NAME(1, BFloat16, GatedAdd); class CUDA_MS_OP_TYPED_CLASS_NAME(1, float, CausalConvWithState); class CUDA_MS_OP_TYPED_CLASS_NAME(1, MLFloat16, CausalConvWithState); +class CUDA_MS_OP_TYPED_CLASS_NAME(1, BFloat16, CausalConvWithState); #if !defined(DISABLE_GENERATION_OPS) class CUDA_MS_OP_CLASS_NAME(1, Sampling); #endif @@ -441,6 +443,7 @@ Status RegisterCudaContribKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, @@ -452,6 +455,7 @@ Status RegisterCudaContribKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, #if !defined(DISABLE_GENERATION_OPS) BuildKernelCreateInfo, #endif diff --git a/onnxruntime/contrib_ops/webgpu/bert/gated_add.cc b/onnxruntime/contrib_ops/webgpu/bert/gated_add.cc new file mode 100644 index 0000000000000..1202fe55857a6 --- /dev/null +++ b/onnxruntime/contrib_ops/webgpu/bert/gated_add.cc @@ -0,0 +1,78 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "contrib_ops/webgpu/bert/gated_add.h" + +#include "core/providers/webgpu/shader_helper.h" +#include "core/providers/webgpu/webgpu_supported_types.h" +#include "contrib_ops/webgpu/webgpu_contrib_kernels.h" + +namespace onnxruntime { +namespace contrib { +namespace webgpu { + +ONNX_OPERATOR_KERNEL_EX( + GatedAdd, + kMSDomain, + 1, + kWebGpuExecutionProvider, + (*KernelDefBuilder::Create()) + .TypeConstraint("T", WebGpuSupportedFloatTypes()), + GatedAdd); + +Status GatedAddProgram::GenerateShaderCode(ShaderHelper& shader) const { + const auto& x = shader.AddInput("x", ShaderUsage::UseUniform); + const auto& y = shader.AddInput("y", ShaderUsage::UseUniform); + const auto& gate = shader.AddInput("gate", ShaderUsage::UseUniform); + const auto& output = shader.AddOutput("output", ShaderUsage::UseUniform); + + shader.MainFunctionBody() << shader.GuardAgainstOutOfBoundsWorkgroupSizes("uniforms.output_size") + << " let gate_idx = global_idx / uniforms.hidden_size;\n" + << " let gate_value = " << gate.GetByOffset("gate_idx") << ";\n" + << " let value = " << x.GetByOffset("global_idx") + << " + (" << y.GetByOffset("global_idx") << " * gate_value);\n" + << " " << output.SetByOffset("global_idx", "value"); + + return Status::OK(); +} + +Status GatedAdd::ComputeInternal(ComputeContext& context) const { + const auto* x = context.Input(0); + const auto* y = context.Input(1); + const auto* gate = context.Input(2); + const TensorShape& shape = x->Shape(); + + ORT_RETURN_IF_NOT(shape.NumDimensions() >= 1, "X must have rank >= 1"); + ORT_RETURN_IF_NOT(y->Shape() == shape, "Y must have the same shape as X"); + ORT_RETURN_IF_NOT(gate->Shape().NumDimensions() == shape.NumDimensions(), + "gate must have the same rank as X"); + + const size_t last_axis = shape.NumDimensions() - 1; + const int64_t hidden_size = shape[last_axis]; + ORT_RETURN_IF_NOT(hidden_size > 0, "X last dimension must be positive"); + ORT_RETURN_IF_NOT(gate->Shape()[last_axis] == 1, "gate last dimension must be 1"); + for (size_t axis = 0; axis < last_axis; ++axis) { + ORT_RETURN_IF_NOT(gate->Shape()[axis] == shape[axis], + "gate dimension ", axis, " must match X"); + } + + auto* output = context.Output(0, shape); + const int64_t output_size = shape.Size(); + if (output_size == 0) { + return Status::OK(); + } + + GatedAddProgram program{}; + program.AddInputs({{x, ProgramTensorMetadataDependency::Type}, + {y, ProgramTensorMetadataDependency::Type}, + {gate, ProgramTensorMetadataDependency::Type}}) + .AddOutput({output, ProgramTensorMetadataDependency::None}) + .SetDispatchGroupSize((onnxruntime::narrow(output_size) + WORKGROUP_SIZE - 1) / WORKGROUP_SIZE) + .AddUniformVariables({{onnxruntime::narrow(output_size)}, + {onnxruntime::narrow(hidden_size)}}); + return context.RunProgram(program); +} + +} // namespace webgpu +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/webgpu/bert/gated_add.h b/onnxruntime/contrib_ops/webgpu/bert/gated_add.h new file mode 100644 index 0000000000000..1bb06237a39ff --- /dev/null +++ b/onnxruntime/contrib_ops/webgpu/bert/gated_add.h @@ -0,0 +1,33 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include "core/providers/webgpu/program.h" +#include "core/providers/webgpu/webgpu_kernel.h" + +namespace onnxruntime { +namespace contrib { +namespace webgpu { + +using namespace onnxruntime::webgpu; +using onnxruntime::webgpu::ComputeContext; + +// output = X + Y * gate, with gate broadcast across the last dimension. +class GatedAddProgram final : public Program { + public: + GatedAddProgram() : Program{"GatedAdd"} {} + Status GenerateShaderCode(ShaderHelper& sh) const override; + WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES({"output_size", ProgramUniformVariableDataType::Uint32}, + {"hidden_size", ProgramUniformVariableDataType::Uint32}); +}; + +class GatedAdd final : public WebGpuKernel { + public: + GatedAdd(const OpKernelInfo& info) : WebGpuKernel(info) {} + Status ComputeInternal(ComputeContext& context) const override; +}; + +} // namespace webgpu +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/webgpu/bert/linear_attention_gates.cc b/onnxruntime/contrib_ops/webgpu/bert/linear_attention_gates.cc new file mode 100644 index 0000000000000..399e64a03ff52 --- /dev/null +++ b/onnxruntime/contrib_ops/webgpu/bert/linear_attention_gates.cc @@ -0,0 +1,216 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "contrib_ops/webgpu/bert/linear_attention_gates.h" + +#include "core/providers/webgpu/shader_helper.h" +#include "core/providers/webgpu/webgpu_supported_types.h" +#include "contrib_ops/webgpu/webgpu_contrib_kernels.h" + +namespace onnxruntime { +namespace contrib { +namespace webgpu { + +ONNX_OPERATOR_KERNEL_EX( + LinearAttentionGate, + kMSDomain, + 1, + kWebGpuExecutionProvider, + (*KernelDefBuilder::Create()) + .TypeConstraint("T", WebGpuSupportedFloatTypes()) + .TypeConstraint("TF", DataTypeImpl::GetTensorType()), + LinearAttentionGate); + +ONNX_OPERATOR_KERNEL_EX( + GatedRMSNorm, + kMSDomain, + 1, + kWebGpuExecutionProvider, + (*KernelDefBuilder::Create()) + .TypeConstraint("T", WebGpuSupportedFloatTypes()), + GatedRMSNorm); + +Status LinearAttentionGateProgram::GenerateShaderCode(ShaderHelper& shader) const { + const auto& a = shader.AddInput("a", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); + const auto& dt_bias = shader.AddInput("dt_bias", ShaderUsage::UseUniform); + const auto& decay_scale = shader.AddInput("decay_scale", ShaderUsage::UseUniform); + const ShaderVariableHelper* b = nullptr; + if (has_b_) { + b = &shader.AddInput("b", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); + } + const auto& decay = shader.AddOutput("decay", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); + const ShaderVariableHelper* beta = nullptr; + if (has_beta_) { + beta = &shader.AddOutput("beta", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); + } + + shader.AdditionalImplementation() + << "fn la_softplus(x: f32) -> f32 {\n" + << " if (x > 0.0) {\n" + << " return x + log(exp(-x) + 1.0);\n" + << " }\n" + << " return log(exp(x) + 1.0);\n" + << "}\n" + << "fn la_sigmoid(x: f32) -> f32 {\n" + << " if (x > 0.0) {\n" + << " return 1.0 / (1.0 + exp(-x));\n" + << " }\n" + << " let e = exp(x);\n" + << " return e / (1.0 + e);\n" + << "}\n"; + + shader.MainFunctionBody() << shader.GuardAgainstOutOfBoundsWorkgroupSizes("uniforms.output_size") + << " let h = global_idx % uniforms.num_heads;\n" + << " let biased = f32(" << a.GetByOffset("global_idx") << ") + f32(" + << dt_bias.GetByOffset("h") << ");\n" + << " " << decay.SetByOffset("global_idx", "decay_element_t(f32(" + decay_scale.GetByOffset("h") + ") * la_softplus(biased))") + << "\n"; + if (has_beta_) { + shader.MainFunctionBody() << " " + << beta->SetByOffset("global_idx", + "beta_element_t(la_sigmoid(f32(" + b->GetByOffset("global_idx") + ")))") + << "\n"; + } + + return Status::OK(); +} + +Status LinearAttentionGate::ComputeInternal(ComputeContext& context) const { + const auto* a = context.Input(0); + const auto* dt_bias = context.Input(1); + const auto* decay_scale = context.Input(2); + const auto* b = context.Input(3); // optional + + const auto& a_shape = a->Shape(); + ORT_RETURN_IF_NOT(a_shape.NumDimensions() >= 1, "a must have rank >= 1"); + const int64_t num_heads = a_shape[a_shape.NumDimensions() - 1]; + ORT_RETURN_IF_NOT(num_heads > 0, "a last dimension must be positive"); + + ORT_RETURN_IF_NOT(dt_bias->Shape().Size() == num_heads, + "dt_bias must have ", num_heads, " elements, got ", dt_bias->Shape().Size()); + ORT_RETURN_IF_NOT(decay_scale->Shape().Size() == num_heads, + "decay_scale must have ", num_heads, " elements, got ", decay_scale->Shape().Size()); + + auto* decay = context.Output(0, a_shape); + auto* beta = context.Output(1, a_shape); + + if (beta != nullptr) { + ORT_RETURN_IF_NOT(b != nullptr, "The b input is required when the beta output is requested"); + ORT_RETURN_IF_NOT(b->Shape() == a_shape, "b must have the same shape as a"); + } + + const int64_t output_size = a_shape.Size(); + if (output_size == 0) { + return Status::OK(); + } + + LinearAttentionGateProgram program{b != nullptr, beta != nullptr}; + program.CacheHint(b != nullptr, beta != nullptr) + .AddInputs({{a, ProgramTensorMetadataDependency::Type}, + {dt_bias, ProgramTensorMetadataDependency::None}, + {decay_scale, ProgramTensorMetadataDependency::None}}); + if (b != nullptr) { + program.AddInput({b, ProgramTensorMetadataDependency::Type}); + } + program.AddOutput({decay, ProgramTensorMetadataDependency::None}); + if (beta != nullptr) { + program.AddOutput({beta, ProgramTensorMetadataDependency::None}); + } + program.SetDispatchGroupSize((onnxruntime::narrow(output_size) + WORKGROUP_SIZE - 1) / WORKGROUP_SIZE) + .AddUniformVariables({{onnxruntime::narrow(output_size)}, + {onnxruntime::narrow(num_heads)}}); + + return context.RunProgram(program); +} + +Status GatedRMSNormProgram::GenerateShaderCode(ShaderHelper& shader) const { + const auto& input = shader.AddInput("input", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); + const auto& scale = shader.AddInput("scale", ShaderUsage::UseUniform); + const auto& gate = shader.AddInput("gate", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); + const auto& output = shader.AddOutput("output", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); + + shader.AdditionalImplementation() + << "fn stable_sigmoid(x: f32) -> f32 {\n" + << " if (x > 0.0) {\n" + << " return 1.0 / (1.0 + exp(-x));\n" + << " }\n" + << " let e = exp(x);\n" + << " return e / (1.0 + e);\n" + << "}\n" + << "var row_sum_sq : array;\n"; + + shader.MainFunctionBody() + << " let row = workgroup_idx;\n" + << " let base = row * uniforms.norm_size;\n" + << " var sum_sq = 0.0;\n" + << " for (var i = local_idx; i < uniforms.norm_size; i += workgroup_size_x) {\n" + << " let v = f32(" << input.GetByOffset("base + i") << ");\n" + << " sum_sq += v * v;\n" + << " }\n" + << " row_sum_sq[local_idx] = sum_sq;\n" + << " workgroupBarrier();\n" + << " var reduce_size = workgroup_size_x;\n" + << " for (var curr_size = reduce_size >> 1u; curr_size > 0u; curr_size = reduce_size >> 1u) {\n" + << " reduce_size = curr_size + (reduce_size & 1u);\n" + << " if (local_idx < curr_size) {\n" + << " row_sum_sq[local_idx] += row_sum_sq[local_idx + reduce_size];\n" + << " }\n" + << " workgroupBarrier();\n" + << " }\n" + << " let inv_rms = inverseSqrt(row_sum_sq[0] / f32(uniforms.norm_size) + uniforms.epsilon);\n" + << " for (var i = local_idx; i < uniforms.norm_size; i += workgroup_size_x) {\n" + << " let z = f32(" << gate.GetByOffset("base + i") << ");\n" + << " let normalized = f32(" << input.GetByOffset("base + i") << ") * inv_rms * f32(" + << scale.GetByOffset("i") << ");\n" + << " " << output.SetByOffset("base + i", "output_element_t(normalized * (z * stable_sigmoid(z)))") << "\n" + << " }\n"; + + return Status::OK(); +} + +GatedRMSNorm::GatedRMSNorm(const OpKernelInfo& info) : WebGpuKernel(info) { + epsilon_ = info.GetAttrOrDefault("epsilon", 1e-5f); +} + +Status GatedRMSNorm::ComputeInternal(ComputeContext& context) const { + const auto* input = context.Input(0); + const auto* scale = context.Input(1); + const auto* gate = context.Input(2); + + const auto& shape = input->Shape(); + ORT_RETURN_IF_NOT(shape.NumDimensions() >= 1, "X must have rank >= 1"); + ORT_RETURN_IF_NOT(gate->Shape() == shape, "gate must have the same shape as X"); + + const int64_t norm_size = scale->Shape().Size(); + ORT_RETURN_IF_NOT(norm_size > 0, "scale must not be empty"); + const int64_t last_dim = shape[shape.NumDimensions() - 1]; + ORT_RETURN_IF_NOT(last_dim % norm_size == 0, + "X last dimension (", last_dim, ") must be a multiple of the scale length (", + norm_size, ")"); + + auto* output = context.Output(0, shape); + const int64_t total_size = shape.Size(); + if (total_size == 0) { + return Status::OK(); + } + const int64_t num_rows = total_size / norm_size; + + const uint32_t workgroup_size = norm_size <= 64 ? 64 + : norm_size <= 128 ? 128 + : 256; + + GatedRMSNormProgram program{}; + program.AddInputs({{input, ProgramTensorMetadataDependency::Type}, + {scale, ProgramTensorMetadataDependency::Type}, + {gate, ProgramTensorMetadataDependency::Type}}) + .AddOutput({output, ProgramTensorMetadataDependency::None}) + .SetDispatchGroupSize(onnxruntime::narrow(num_rows)) + .SetWorkgroupSize(workgroup_size) + .AddUniformVariables({{onnxruntime::narrow(norm_size)}, + {epsilon_}}); + return context.RunProgram(program); +} + +} // namespace webgpu +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/webgpu/bert/linear_attention_gates.h b/onnxruntime/contrib_ops/webgpu/bert/linear_attention_gates.h new file mode 100644 index 0000000000000..f4910cb45602d --- /dev/null +++ b/onnxruntime/contrib_ops/webgpu/bert/linear_attention_gates.h @@ -0,0 +1,55 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include "core/providers/webgpu/program.h" +#include "core/providers/webgpu/webgpu_kernel.h" + +namespace onnxruntime { +namespace contrib { +namespace webgpu { + +using namespace onnxruntime::webgpu; +using onnxruntime::webgpu::ComputeContext; + +// decay = decay_scale * Softplus(a + dt_bias), beta = Sigmoid(b). +class LinearAttentionGateProgram final : public Program { + public: + LinearAttentionGateProgram(bool has_b, bool has_beta) : Program{"LinearAttentionGate"}, has_b_(has_b), has_beta_(has_beta) {} + Status GenerateShaderCode(ShaderHelper& sh) const override; + WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES({"output_size", ProgramUniformVariableDataType::Uint32}, + {"num_heads", ProgramUniformVariableDataType::Uint32}); + + private: + bool has_b_; + bool has_beta_; +}; + +class LinearAttentionGate final : public WebGpuKernel { + public: + LinearAttentionGate(const OpKernelInfo& info) : WebGpuKernel(info) {} + Status ComputeInternal(ComputeContext& context) const override; +}; + +// Y = X * rsqrt(mean(X^2) + epsilon) * scale * SiLU(gate). +class GatedRMSNormProgram final : public Program { + public: + GatedRMSNormProgram() : Program{"GatedRMSNorm"} {} + Status GenerateShaderCode(ShaderHelper& sh) const override; + WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES({"norm_size", ProgramUniformVariableDataType::Uint32}, + {"epsilon", ProgramUniformVariableDataType::Float32}); +}; + +class GatedRMSNorm final : public WebGpuKernel { + public: + GatedRMSNorm(const OpKernelInfo& info); + Status ComputeInternal(ComputeContext& context) const override; + + private: + float epsilon_; +}; + +} // namespace webgpu +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/webgpu/webgpu_contrib_kernels.cc b/onnxruntime/contrib_ops/webgpu/webgpu_contrib_kernels.cc index e827f2042e428..6d1e283eae13d 100644 --- a/onnxruntime/contrib_ops/webgpu/webgpu_contrib_kernels.cc +++ b/onnxruntime/contrib_ops/webgpu/webgpu_contrib_kernels.cc @@ -3,8 +3,10 @@ #include "contrib_ops/webgpu/webgpu_contrib_kernels.h" #include "contrib_ops/webgpu/bert/causal_conv_with_state.h" +#include "contrib_ops/webgpu/bert/gated_add.h" #include "contrib_ops/webgpu/bert/group_query_attention.h" #include "contrib_ops/webgpu/bert/linear_attention.h" +#include "contrib_ops/webgpu/bert/linear_attention_gates.h" #include "contrib_ops/webgpu/bert/paged_attention.h" #include "core/framework/op_kernel.h" @@ -29,8 +31,11 @@ static const BuildKernelCreateInfoFn build_kernel_create_info_function_table[] = BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, diff --git a/onnxruntime/core/providers/cuda/cuda_execution_provider.cc b/onnxruntime/core/providers/cuda/cuda_execution_provider.cc index d2515c465e16a..f624c3f66b73c 100755 --- a/onnxruntime/core/providers/cuda/cuda_execution_provider.cc +++ b/onnxruntime/core/providers/cuda/cuda_execution_provider.cc @@ -1665,6 +1665,7 @@ class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDom // Opset 22. class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, float, LpNormalization); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, MLFloat16, LpNormalization); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, BFloat16, LpNormalization); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, float, AveragePool); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, double, AveragePool); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, MLFloat16, AveragePool); @@ -2956,6 +2957,7 @@ static Status RegisterCudaKernels(KernelRegistry& kernel_registry) { // Opset 22 BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, diff --git a/onnxruntime/core/providers/cuda/nn/lp_norm.cc b/onnxruntime/core/providers/cuda/nn/lp_norm.cc index 6a765941a298e..1d5bd4ff9d250 100644 --- a/onnxruntime/core/providers/cuda/nn/lp_norm.cc +++ b/onnxruntime/core/providers/cuda/nn/lp_norm.cc @@ -35,6 +35,7 @@ REGISTER_LPNORM_VERSIONED_KERNEL(MLFloat16, 1, 21) REGISTER_LPNORM_KERNEL(float, 22) REGISTER_LPNORM_KERNEL(MLFloat16, 22) +REGISTER_LPNORM_KERNEL(BFloat16, 22) template Status LpNorm::ComputeInternal(OpKernelContext* context) const { diff --git a/onnxruntime/core/providers/cuda/nn/lp_norm_impl.cu b/onnxruntime/core/providers/cuda/nn/lp_norm_impl.cu index db36861546960..1d83625119494 100644 --- a/onnxruntime/core/providers/cuda/nn/lp_norm_impl.cu +++ b/onnxruntime/core/providers/cuda/nn/lp_norm_impl.cu @@ -124,6 +124,7 @@ void LpNormImpl( // Explicit instantiations. template void LpNormImpl(cudaStream_t, const float*, float*, int64_t, int64_t, int64_t, int); template void LpNormImpl(cudaStream_t, const half*, half*, int64_t, int64_t, int64_t, int); +template void LpNormImpl(cudaStream_t, const BFloat16*, BFloat16*, int64_t, int64_t, int64_t, int); } // namespace cuda } // namespace onnxruntime diff --git a/onnxruntime/test/contrib_ops/causal_conv_with_state_op_test.cc b/onnxruntime/test/contrib_ops/causal_conv_with_state_op_test.cc index 08a6e0f1e3fce..e0c9f6b4f206c 100644 --- a/onnxruntime/test/contrib_ops/causal_conv_with_state_op_test.cc +++ b/onnxruntime/test/contrib_ops/causal_conv_with_state_op_test.cc @@ -840,6 +840,68 @@ TEST(CausalConvWithStateTest, StateWindow_DecodeGenericKWithPastState) { RunCausalConvStateWindowTest(/*batch_size=*/2, /*channels=*/8, /*input_length=*/1, /*kernel_size=*/7, /*window=*/3, /*with_past_state=*/true); } + +// BFloat16 is CUDA-only for this op (CPU/WebGPU only register float/float16). +TEST(CausalConvWithStateTest, BFloat16_Cuda) { + auto ep = DefaultCudaExecutionProvider(); + if (!ep) { + GTEST_SKIP() << "CUDA execution provider not available"; + return; + } + if (!CudaHasBF16Support()) { + GTEST_SKIP() << "CUDA device does not support BFloat16."; + return; + } + + const int batch_size = 2; + const int channels = 4; + const int input_length = 6; + const int kernel_size = 3; + const int state_length = kernel_size - 1; + const std::string activation = "silu"; + + std::vector input_data(static_cast(batch_size) * channels * input_length); + for (size_t i = 0; i < input_data.size(); ++i) { + input_data[i] = 0.3f * std::sin(static_cast(i) * 0.41f); + } + std::vector weight_data(static_cast(channels) * kernel_size); + for (size_t i = 0; i < weight_data.size(); ++i) { + weight_data[i] = 0.2f * std::cos(static_cast(i) * 0.17f); + } + std::vector bias_data(channels); + for (int i = 0; i < channels; ++i) { + bias_data[i] = 0.01f * static_cast(i); + } + std::vector conv_state_data(static_cast(batch_size) * channels * state_length); + for (size_t i = 0; i < conv_state_data.size(); ++i) { + conv_state_data[i] = 0.25f * std::sin(static_cast(i) * 0.29f); + } + + std::vector expected_output, expected_state; + CausalConvWithStateReference( + input_data, weight_data, &bias_data, &conv_state_data, + expected_output, expected_state, + batch_size, channels, input_length, kernel_size, activation); + + std::vector input_shape = {batch_size, channels, input_length}; + std::vector weight_shape = {channels, 1, kernel_size}; + std::vector bias_shape = {channels}; + std::vector state_shape = {batch_size, channels, state_length}; + std::vector output_shape = {batch_size, channels, input_length}; + + OpTester test("CausalConvWithState", 1, onnxruntime::kMSDomain); + test.AddAttribute("activation", activation); + test.AddInput("input", input_shape, ToBFloat16(input_data)); + test.AddInput("weight", weight_shape, ToBFloat16(weight_data)); + test.AddInput("bias", bias_shape, ToBFloat16(bias_data)); + test.AddInput("past_state", state_shape, ToBFloat16(conv_state_data)); + test.AddOutput("output", output_shape, ToBFloat16(expected_output), false, 0.02f, 0.0f); + test.AddOutput("present_state", state_shape, ToBFloat16(expected_state), false, 0.02f, 0.0f); + + std::vector> execution_providers; + execution_providers.push_back(std::move(ep)); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); +} #endif // USE_CUDA } // namespace test diff --git a/onnxruntime/test/contrib_ops/gated_add_op_test.cc b/onnxruntime/test/contrib_ops/gated_add_op_test.cc index e33a6ed095035..d3a978c076a5a 100644 --- a/onnxruntime/test/contrib_ops/gated_add_op_test.cc +++ b/onnxruntime/test/contrib_ops/gated_add_op_test.cc @@ -16,6 +16,20 @@ namespace test { namespace { +// Collects every EP that has a GatedAdd kernel available in this build, so the fused-op tests +// exercise all of them. +std::vector> AvailableGatedOpExecutionProviders() { + std::vector> eps; + eps.push_back(DefaultCpuExecutionProvider()); + if (auto cuda_ep = DefaultCudaExecutionProvider()) { + eps.push_back(std::move(cuda_ep)); + } + if (auto webgpu_ep = DefaultWebGpuExecutionProvider()) { + eps.push_back(std::move(webgpu_ep)); + } + return eps; +} + template float RoundToType(float value) { if constexpr (std::is_same_v) { @@ -38,13 +52,22 @@ std::vector ToTensorType(const std::vector& data) { } } +// BFloat16 is currently CUDA-only; other types run on every available EP. template -void RunGatedAddTest(const std::vector& input_dims) { - auto cuda_ep = DefaultCudaExecutionProvider(); - if (!cuda_ep) { - GTEST_SKIP() << "CUDA EP not available"; +std::vector> ExecutionProvidersForType() { + if constexpr (std::is_same_v) { + std::vector> eps; + if (auto cuda_ep = DefaultCudaExecutionProvider()) { + eps.push_back(std::move(cuda_ep)); + } + return eps; + } else { + return AvailableGatedOpExecutionProviders(); } +} +template +void RunGatedAddTest(const std::vector& input_dims) { ASSERT_FALSE(input_dims.empty()); int64_t rows = 1; for (size_t axis = 0; axis + 1 < input_dims.size(); ++axis) { @@ -72,18 +95,29 @@ void RunGatedAddTest(const std::vector& input_dims) { std::vector gate_dims = input_dims; gate_dims.back() = 1; - OpTester tester("GatedAdd", 1, onnxruntime::kMSDomain); - tester.AddInput("X", input_dims, ToTensorType(x)); - tester.AddInput("Y", input_dims, ToTensorType(y)); - tester.AddInput("gate", gate_dims, ToTensorType(gate)); - tester.AddOutput("output", input_dims, ToTensorType(expected)); - if constexpr (!std::is_same_v) { - tester.SetOutputTolerance(0.0f, 0.0f); + + auto execution_providers = ExecutionProvidersForType(); + if (execution_providers.empty()) { + GTEST_SKIP() << "No execution provider available for this type"; } - std::vector> execution_providers; - execution_providers.push_back(std::move(cuda_ep)); - tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); + for (auto& ep : execution_providers) { + SCOPED_TRACE("EP: " + ep->Type()); + OpTester tester("GatedAdd", 1, onnxruntime::kMSDomain); + tester.AddInput("X", input_dims, ToTensorType(x)); + tester.AddInput("Y", input_dims, ToTensorType(y)); + tester.AddInput("gate", gate_dims, ToTensorType(gate)); + if constexpr (std::is_same_v) { + // Allows valid WebGPU FP16 rounding differences near cancellation. + tester.AddOutput("output", input_dims, ToTensorType(expected), false, 0.001f, 0.01f); + } else { + tester.AddOutput("output", input_dims, ToTensorType(expected)); + } + + std::vector> providers; + providers.push_back(std::move(ep)); + tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &providers); + } } } // namespace @@ -101,21 +135,20 @@ TEST(ContribOpGatedAddTest, EmptyOuterDimension) { } TEST(ContribOpGatedAddTest, ZeroHiddenDimension) { - auto cuda_ep = DefaultCudaExecutionProvider(); - if (!cuda_ep) { - GTEST_SKIP() << "CUDA EP not available"; + auto execution_providers = AvailableGatedOpExecutionProviders(); + for (auto& ep : execution_providers) { + SCOPED_TRACE("EP: " + ep->Type()); + OpTester tester("GatedAdd", 1, onnxruntime::kMSDomain); + tester.AddInput("X", {2, 3, 0}, {}); + tester.AddInput("Y", {2, 3, 0}, {}); + tester.AddInput("gate", {2, 3, 1}, std::vector(6)); + tester.AddOutput("output", {2, 3, 0}, {}); + + std::vector> providers; + providers.push_back(std::move(ep)); + tester.Run(OpTester::ExpectResult::kExpectFailure, "X last dimension must be positive", + {}, nullptr, &providers); } - - OpTester tester("GatedAdd", 1, onnxruntime::kMSDomain); - tester.AddInput("X", {2, 3, 0}, {}); - tester.AddInput("Y", {2, 3, 0}, {}); - tester.AddInput("gate", {2, 3, 1}, std::vector(6)); - tester.AddOutput("output", {2, 3, 0}, {}); - - std::vector> execution_providers; - execution_providers.push_back(std::move(cuda_ep)); - tester.Run(OpTester::ExpectResult::kExpectFailure, "X last dimension must be positive", - {}, nullptr, &execution_providers); } TEST(ContribOpGatedAddTest, Float16) { @@ -130,22 +163,21 @@ TEST(ContribOpGatedAddTest, BFloat16) { } TEST(ContribOpGatedAddTest, MismatchedYShape) { - auto cuda_ep = DefaultCudaExecutionProvider(); - if (!cuda_ep) { - GTEST_SKIP() << "CUDA EP not available"; + auto execution_providers = AvailableGatedOpExecutionProviders(); + for (auto& ep : execution_providers) { + SCOPED_TRACE("EP: " + ep->Type()); + OpTester tester("GatedAdd", 1, onnxruntime::kMSDomain); + tester.AddInput("X", {1, 2, 3}, std::vector(6)); + tester.AddInput("Y", {1, 2, 4}, std::vector(8)); + tester.AddInput("gate", {1, 2, 1}, std::vector(2)); + tester.AddOutput("output", {1, 2, 3}, std::vector(6)); + + std::vector> providers; + providers.push_back(std::move(ep)); + tester.Run(OpTester::ExpectResult::kExpectFailure, "Y must have the same shape as X", + {}, nullptr, &providers); } - - OpTester tester("GatedAdd", 1, onnxruntime::kMSDomain); - tester.AddInput("X", {1, 2, 3}, std::vector(6)); - tester.AddInput("Y", {1, 2, 4}, std::vector(8)); - tester.AddInput("gate", {1, 2, 1}, std::vector(2)); - tester.AddOutput("output", {1, 2, 3}, std::vector(6)); - - std::vector> execution_providers; - execution_providers.push_back(std::move(cuda_ep)); - tester.Run(OpTester::ExpectResult::kExpectFailure, "Y must have the same shape as X", - {}, nullptr, &execution_providers); } } // namespace test -} // namespace onnxruntime \ No newline at end of file +} // namespace onnxruntime diff --git a/onnxruntime/test/contrib_ops/linear_attention_gates_op_test.cc b/onnxruntime/test/contrib_ops/linear_attention_gates_op_test.cc index 41def5b5d06d0..cfea079e36c32 100644 --- a/onnxruntime/test/contrib_ops/linear_attention_gates_op_test.cc +++ b/onnxruntime/test/contrib_ops/linear_attention_gates_op_test.cc @@ -4,6 +4,7 @@ // Tests for the two fused linear-attention gate ops: LinearAttentionGate and GatedRMSNorm. // Both replace float32 elementwise chains that an exporter emits around the LinearAttention op, // so the references here are the same float32 formulas ORT's Softplus/Sigmoid/RMSNorm kernels use. +// The tests run against every EP (CPU, CUDA, WebGPU) that has these ops registered. #include #include @@ -21,6 +22,32 @@ namespace test { namespace { +std::vector> AvailableGatedOpExecutionProviders() { + std::vector> eps; + eps.push_back(DefaultCpuExecutionProvider()); + if (auto cuda_ep = DefaultCudaExecutionProvider()) { + eps.push_back(std::move(cuda_ep)); + } + if (auto webgpu_ep = DefaultWebGpuExecutionProvider()) { + eps.push_back(std::move(webgpu_ep)); + } + return eps; +} + +// BFloat16 is currently CUDA-only; other types run on every available EP. +template +std::vector> ExecutionProvidersForType() { + if constexpr (std::is_same_v) { + std::vector> eps; + if (auto cuda_ep = DefaultCudaExecutionProvider()) { + eps.push_back(std::move(cuda_ep)); + } + return eps; + } else { + return AvailableGatedOpExecutionProviders(); + } +} + float SigmoidRef(float x) { return x > 0.0f ? 1.0f / (1.0f + std::exp(-x)) : 1.0f - 1.0f / (1.0f + std::exp(x)); } @@ -53,9 +80,9 @@ std::vector ToTensorType(const std::vector& data) { template void RunLinearAttentionGateTest(int batch_size, int seq_length, int num_heads, bool with_beta, float tolerance) { - auto cuda_ep = DefaultCudaExecutionProvider(); - if (!cuda_ep) { - GTEST_SKIP() << "CUDA EP not available"; + auto execution_providers = ExecutionProvidersForType(); + if (execution_providers.empty()) { + GTEST_SKIP() << "No execution provider available for this type"; } const size_t count = static_cast(batch_size) * seq_length * num_heads; @@ -75,31 +102,34 @@ void RunLinearAttentionGateTest(int batch_size, int seq_length, int num_heads, b const std::vector dims = {batch_size, seq_length, num_heads}; const std::vector param_dims = {num_heads}; - OpTester tester("LinearAttentionGate", 1, onnxruntime::kMSDomain); - tester.AddInput("a", dims, ToTensorType(a)); - tester.AddInput("dt_bias", param_dims, dt_bias); - tester.AddInput("decay_scale", param_dims, decay_scale); - if (with_beta) { - tester.AddInput("b", dims, ToTensorType(b)); - } else { - tester.AddOptionalInputEdge(); - } - tester.AddOutput("decay", dims, ToTensorType(expected_decay), false, tolerance, tolerance); - if (with_beta) { - tester.AddOutput("beta", dims, ToTensorType(expected_beta), false, tolerance, tolerance); - } + for (auto& ep : execution_providers) { + SCOPED_TRACE("EP: " + ep->Type()); + OpTester tester("LinearAttentionGate", 1, onnxruntime::kMSDomain); + tester.AddInput("a", dims, ToTensorType(a)); + tester.AddInput("dt_bias", param_dims, dt_bias); + tester.AddInput("decay_scale", param_dims, decay_scale); + if (with_beta) { + tester.AddInput("b", dims, ToTensorType(b)); + } else { + tester.AddOptionalInputEdge(); + } + tester.AddOutput("decay", dims, ToTensorType(expected_decay), false, tolerance, tolerance); + if (with_beta) { + tester.AddOutput("beta", dims, ToTensorType(expected_beta), false, tolerance, tolerance); + } - std::vector> execution_providers; - execution_providers.push_back(std::move(cuda_ep)); - tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); + std::vector> providers; + providers.push_back(std::move(ep)); + tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &providers); + } } template void RunGatedRMSNormTest(int batch_size, int seq_length, int num_heads, int head_dim, float epsilon, float tolerance) { - auto cuda_ep = DefaultCudaExecutionProvider(); - if (!cuda_ep) { - GTEST_SKIP() << "CUDA EP not available"; + auto execution_providers = ExecutionProvidersForType(); + if (execution_providers.empty()) { + GTEST_SKIP() << "No execution provider available for this type"; } const int hidden = num_heads * head_dim; @@ -126,16 +156,19 @@ void RunGatedRMSNormTest(int batch_size, int seq_length, int num_heads, int head const std::vector dims = {batch_size, seq_length, hidden}; const std::vector scale_dims = {head_dim}; - OpTester tester("GatedRMSNorm", 1, onnxruntime::kMSDomain); - tester.AddAttribute("epsilon", epsilon); - tester.AddInput("X", dims, ToTensorType(x)); - tester.AddInput("scale", scale_dims, ToTensorType(scale)); - tester.AddInput("gate", dims, ToTensorType(gate)); - tester.AddOutput("Y", dims, ToTensorType(expected), false, tolerance, tolerance); - - std::vector> execution_providers; - execution_providers.push_back(std::move(cuda_ep)); - tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); + for (auto& ep : execution_providers) { + SCOPED_TRACE("EP: " + ep->Type()); + OpTester tester("GatedRMSNorm", 1, onnxruntime::kMSDomain); + tester.AddAttribute("epsilon", epsilon); + tester.AddInput("X", dims, ToTensorType(x)); + tester.AddInput("scale", scale_dims, ToTensorType(scale)); + tester.AddInput("gate", dims, ToTensorType(gate)); + tester.AddOutput("Y", dims, ToTensorType(expected), false, tolerance, tolerance); + + std::vector> providers; + providers.push_back(std::move(ep)); + tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &providers); + } } } // namespace @@ -174,10 +207,7 @@ TEST(ContribOpLinearAttentionGateTest, BFloat16_SpeculativeDecodeTile) { // Requesting beta without b must be rejected by shape inference, not at execution time. TEST(ContribOpLinearAttentionGateTest, BetaWithoutB_FailsShapeInference) { - auto cuda_ep = DefaultCudaExecutionProvider(); - if (!cuda_ep) { - GTEST_SKIP() << "CUDA EP not available"; - } + auto execution_providers = AvailableGatedOpExecutionProviders(); constexpr int kNumHeads = 8; const std::vector dims = {1, 2, kNumHeads}; @@ -185,19 +215,22 @@ TEST(ContribOpLinearAttentionGateTest, BetaWithoutB_FailsShapeInference) { const std::vector values(static_cast(2 * kNumHeads), 0.5f); const std::vector params(kNumHeads, 0.5f); - OpTester tester("LinearAttentionGate", 1, onnxruntime::kMSDomain); - tester.AddInput("a", dims, values); - tester.AddInput("dt_bias", param_dims, params); - tester.AddInput("decay_scale", param_dims, params); - tester.AddOptionalInputEdge(); - tester.AddOutput("decay", dims, values); - tester.AddOutput("beta", dims, values); - - std::vector> execution_providers; - execution_providers.push_back(std::move(cuda_ep)); - tester.Run(OpTester::ExpectResult::kExpectFailure, - "The b input is required when the beta output is requested", - {}, nullptr, &execution_providers); + for (auto& ep : execution_providers) { + SCOPED_TRACE("EP: " + ep->Type()); + OpTester tester("LinearAttentionGate", 1, onnxruntime::kMSDomain); + tester.AddInput("a", dims, values); + tester.AddInput("dt_bias", param_dims, params); + tester.AddInput("decay_scale", param_dims, params); + tester.AddOptionalInputEdge(); + tester.AddOutput("decay", dims, values); + tester.AddOutput("beta", dims, values); + + std::vector> providers; + providers.push_back(std::move(ep)); + tester.Run(OpTester::ExpectResult::kExpectFailure, + "The b input is required when the beta output is requested", + {}, nullptr, &providers); + } } TEST(ContribOpGatedRMSNormTest, Float_PerHead) { diff --git a/onnxruntime/test/contrib_ops/linear_attention_op_test.cc b/onnxruntime/test/contrib_ops/linear_attention_op_test.cc index 5f20ff14e007b..01b2c3a91e9ee 100644 --- a/onnxruntime/test/contrib_ops/linear_attention_op_test.cc +++ b/onnxruntime/test/contrib_ops/linear_attention_op_test.cc @@ -3,11 +3,14 @@ #include #include +#include #include #include "gtest/gtest.h" #include "core/common/logging/logging.h" #include "core/framework/kernel_registry.h" +#include "test/common/cuda_op_test_utils.h" +#include "test/common/tensor_op_test_utils.h" #include "test/providers/provider_test_utils.h" #include "test/util/include/default_providers.h" @@ -1725,6 +1728,86 @@ TEST(ContribOpLinearAttentionTest, GatedDeltaRule_StateWindow_WiderThanSequenceW RunLinearAttentionStateWindowTest(/*B=*/2, /*q_H=*/2, /*kv_H=*/2, /*n_k=*/1, /*T=*/2, /*dk=*/128, /*dv=*/128, /*W=*/5, /*with_past_state=*/true); } + +// BFloat16 is CUDA-only for this op (CPU/WebGPU only register float/float16). +TEST(ContribOpLinearAttentionTest, BFloat16_Cuda) { + auto ep = DefaultCudaExecutionProvider(); + if (!ep) { + GTEST_SKIP() << "CUDA execution provider not available"; + return; + } + if (!CudaHasBF16Support()) { + GTEST_SKIP() << "CUDA device does not support BFloat16."; + return; + } + + const std::string update_rule = "gated_delta"; + const int batch_size = 2; + const int num_heads = 2; + const int seq_length = 4; + const int head_dim_k = 8; + const int head_dim_v = 8; + const float scale = 1.0f / std::sqrt(static_cast(head_dim_k)); + + auto make_data = [](size_t count, float lo, float hi, uint32_t seed) { + std::vector out(count); + std::mt19937 gen(seed); + std::uniform_real_distribution dist(lo, hi); + for (auto& v : out) v = dist(gen); + return out; + }; + + const size_t qk_count = static_cast(batch_size) * num_heads * seq_length * head_dim_k; + const size_t v_count = static_cast(batch_size) * num_heads * seq_length * head_dim_v; + const size_t state_count = static_cast(batch_size) * num_heads * head_dim_k * head_dim_v; + const size_t decay_count = static_cast(batch_size) * num_heads * seq_length; + const size_t beta_count = decay_count; + + auto query = make_data(qk_count, -0.5f, 0.5f, 1); + auto key = make_data(qk_count, -0.5f, 0.5f, 2); + auto value = make_data(v_count, -0.5f, 0.5f, 3); + auto initial_state = make_data(state_count, -0.1f, 0.1f, 4); + auto decay = make_data(decay_count, -1.0f, 0.0f, 5); + auto beta = make_data(beta_count, 0.1f, 0.9f, 6); + + std::vector expected_output_4d, expected_state; + LinearAttentionReference(update_rule, batch_size, num_heads, seq_length, head_dim_k, head_dim_v, + scale, query, key, value, &initial_state, &decay, &beta, + expected_output_4d, expected_state); + + auto query_3d = PackBHTD_to_BTHD(query, batch_size, num_heads, seq_length, head_dim_k); + auto key_3d = PackBHTD_to_BTHD(key, batch_size, num_heads, seq_length, head_dim_k); + auto value_3d = PackBHTD_to_BTHD(value, batch_size, num_heads, seq_length, head_dim_v); + auto output_3d = PackBHTD_to_BTHD(expected_output_4d, batch_size, num_heads, seq_length, head_dim_v); + auto decay_3d = TransposeBHT_to_BTH(decay, batch_size, num_heads, seq_length); + auto beta_3d = TransposeBHT_to_BTH(beta, batch_size, num_heads, seq_length); + + OpTester tester("LinearAttention", 1, onnxruntime::kMSDomain); + tester.AddAttribute("update_rule", update_rule); + tester.AddAttribute("scale", scale); + tester.AddAttribute("q_num_heads", static_cast(num_heads)); + tester.AddAttribute("kv_num_heads", static_cast(num_heads)); + + std::vector qk_dims = {batch_size, seq_length, num_heads * head_dim_k}; + std::vector v_dims = {batch_size, seq_length, num_heads * head_dim_v}; + std::vector state_dims = {batch_size, num_heads, head_dim_k, head_dim_v}; + std::vector decay_dims = {batch_size, seq_length, num_heads}; + std::vector beta_dims = {batch_size, seq_length, num_heads}; + std::vector out_dims = {batch_size, seq_length, num_heads * head_dim_v}; + + tester.AddInput("query", qk_dims, ToBFloat16(query_3d)); + tester.AddInput("key", qk_dims, ToBFloat16(key_3d)); + tester.AddInput("value", v_dims, ToBFloat16(value_3d)); + tester.AddInput("past_state", state_dims, ToBFloat16(initial_state)); + tester.AddInput("decay", decay_dims, ToBFloat16(decay_3d)); + tester.AddInput("beta", beta_dims, ToBFloat16(beta_3d)); + tester.AddOutput("output", out_dims, ToBFloat16(output_3d), false, 0.02f, 0.0f); + tester.AddOutput("present_state", state_dims, ToBFloat16(expected_state), false, 0.02f, 0.0f); + + std::vector> execution_providers; + execution_providers.push_back(std::move(ep)); + tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); +} #endif // USE_CUDA } // namespace test } // namespace onnxruntime diff --git a/onnxruntime/test/providers/cpu/nn/lp_norm_op_test.cc b/onnxruntime/test/providers/cpu/nn/lp_norm_op_test.cc index 3b4f08aa73379..45b266eac1e81 100644 --- a/onnxruntime/test/providers/cpu/nn/lp_norm_op_test.cc +++ b/onnxruntime/test/providers/cpu/nn/lp_norm_op_test.cc @@ -3,6 +3,7 @@ #include #include "gtest/gtest.h" +#include "test/common/cuda_op_test_utils.h" #include "test/providers/provider_test_utils.h" #include "default_providers.h" #include "core/session/onnxruntime_session_options_config_keys.h" @@ -301,6 +302,49 @@ TEST(LpNormalizationTest, L1Normalization_FP16) { test.Run(so, OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); } +TEST(LpNormalizationTest, L2Normalization_BFloat16_Opset22) { +#ifndef USE_CUDA + GTEST_SKIP() << "BFloat16 tests are only enabled on CUDA builds"; +#else + if (!CudaHasBF16Support()) { + GTEST_SKIP() << "CUDA device does not support BFloat16."; + } + + OpTester test("LpNormalization", 22); + test.AddAttribute("axis", static_cast(-1)); + test.AddAttribute("p", static_cast(2)); + + constexpr int64_t kRows = 2; + constexpr int64_t kCols = 8; + std::vector input_f = { + 1.0f, 2.0f, 3.0f, 4.0f, 2.0f, 1.0f, 0.5f, 0.25f, + 4.0f, 3.0f, 2.0f, 1.0f, 1.0f, 2.0f, 4.0f, 8.0f}; + + std::vector expected_f(kRows * kCols); + for (int64_t r = 0; r < kRows; ++r) { + float sum_sq = 0.0f; + for (int64_t c = 0; c < kCols; ++c) { + const float v = input_f[r * kCols + c]; + sum_sq += v * v; + } + const float norm = std::sqrt(sum_sq); + for (int64_t c = 0; c < kCols; ++c) { + expected_f[r * kCols + c] = input_f[r * kCols + c] / norm; + } + } + + const std::vector dims = {kRows, kCols}; + test.AddInput("input", dims, FloatsToBFloat16s(input_f)); + test.AddOutput("Y", dims, FloatsToBFloat16s(expected_f), false, 0.01f, 0.0f); + + SessionOptions so; + ASSERT_TRUE(so.config_options.AddConfigEntry(kOrtSessionOptionsDisableCPUEPFallback, "1").IsOK()); + std::vector> execution_providers; + execution_providers.push_back(DefaultCudaExecutionProvider()); + test.Run(so, OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); +#endif +} + TEST(LpNormalizationTest, L2Normalization_LastAxis) { // Test normalization along the last axis (axis=-1), which is the most common // use case for LpNormalization in attention patterns (e.g. Qwen3.5 L2-norm).