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