diff --git a/docs/ContribOperators.md b/docs/ContribOperators.md
index c1d12a1d5c..57d15c161b 100644
--- a/docs/ContribOperators.md
+++ b/docs/ContribOperators.md
@@ -127,6 +127,8 @@ This version of the operator has been available since version 1 of the 'com.micr
#### Attributes
+- do_rotary : int
+- Whether to use rotary position embedding. Default value is 0.
- mask_filter_value : float
- The value to be filled in the attention mask. Default value is -10000.0f
- num_heads : int (required)
@@ -2588,6 +2590,8 @@ This version of the operator has been available since version 1 of the 'com.micr
#### Attributes
+- do_rotary : int
+- Whether to use rotary position embedding. Default value is 0.
- mask_filter_value : float
- The value to be filled in the attention mask. Default value is -10000.0f
- num_heads : int (required)
diff --git a/onnxruntime/contrib_ops/cpu/bert/attention_base.cc b/onnxruntime/contrib_ops/cpu/bert/attention_base.cc
index e75f68ea53..e4a8b83159 100644
--- a/onnxruntime/contrib_ops/cpu/bert/attention_base.cc
+++ b/onnxruntime/contrib_ops/cpu/bert/attention_base.cc
@@ -125,6 +125,7 @@ Status AttentionBase::CheckInputs(const TensorShape& input_shape,
}
int64_t past_sequence_length = 0;
+ int64_t original_past_sequence_length = 0;
if (past != nullptr) { // past is optional
if (k_hidden_size != v_hidden_size) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Input 'past' expect k_hidden_size == v_hidden_size");
@@ -157,6 +158,7 @@ Status AttentionBase::CheckInputs(const TensorShape& input_shape,
if (!past_present_share_buffer_) {
past_sequence_length = past_dims[3];
+ original_past_sequence_length = past_sequence_length;
} else {
if (past_seq_len == nullptr || !onnxruntime::IsScalarOr1ElementVector(past_seq_len)) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
@@ -237,6 +239,7 @@ Status AttentionBase::CheckInputs(const TensorShape& input_shape,
output_parameters->batch_size = static_cast(batch_size);
output_parameters->sequence_length = static_cast(sequence_length);
output_parameters->past_sequence_length = static_cast(past_sequence_length);
+ output_parameters->original_past_sequence_length = static_cast(original_past_sequence_length);
output_parameters->kv_sequence_length = static_cast(kv_sequence_length);
output_parameters->total_sequence_length = static_cast(total_sequence_length);
output_parameters->max_sequence_length = static_cast(max_sequence_length);
@@ -248,6 +251,7 @@ Status AttentionBase::CheckInputs(const TensorShape& input_shape,
output_parameters->num_heads = num_heads_;
output_parameters->is_unidirectional = is_unidirectional_;
output_parameters->past_present_share_buffer = (past_present_share_buffer_ != 0 && past != nullptr);
+ output_parameters->do_rotary = do_rotary_;
output_parameters->mask_filter_value = mask_filter_value_;
output_parameters->scale = scale_;
output_parameters->mask_type = mask_type;
diff --git a/onnxruntime/contrib_ops/cpu/bert/attention_base.h b/onnxruntime/contrib_ops/cpu/bert/attention_base.h
index 2e077da285..901f855b25 100644
--- a/onnxruntime/contrib_ops/cpu/bert/attention_base.h
+++ b/onnxruntime/contrib_ops/cpu/bert/attention_base.h
@@ -37,6 +37,7 @@ class AttentionBase {
num_heads_ = static_cast(num_heads);
is_unidirectional_ = info.GetAttrOrDefault("unidirectional", 0) == 1;
+ do_rotary_ = info.GetAttrOrDefault("do_rotary", 0) == 1;
mask_filter_value_ = info.GetAttrOrDefault("mask_filter_value", -10000.0f);
scale_ = info.GetAttrOrDefault("scale", 0.0f);
@@ -70,6 +71,7 @@ class AttentionBase {
std::vector qkv_hidden_sizes_; // Q, K, V hidden sizes parsed from the qkv_hidden_sizes attribute.
bool require_same_hidden_size_; // whether the implementation supports different hidden sizes of Q/K/V.
bool past_present_share_buffer_; // whether or not the past (if used) and present tensor share the same buffer
+ bool do_rotary_; // whether or not to use rotary embeddings
float mask_filter_value_; // the value to be used for filtered out positions
float scale_; // the scale to be used for softmax
};
diff --git a/onnxruntime/contrib_ops/cpu/bert/attention_common.h b/onnxruntime/contrib_ops/cpu/bert/attention_common.h
index bc5ce6e323..adb6805632 100644
--- a/onnxruntime/contrib_ops/cpu/bert/attention_common.h
+++ b/onnxruntime/contrib_ops/cpu/bert/attention_common.h
@@ -38,18 +38,20 @@ enum AttentionKernelType {
struct AttentionParameters {
int batch_size;
int sequence_length;
- int kv_sequence_length; // input sequence length of K or V
- int past_sequence_length; // sequence length in past state of K or V
- int total_sequence_length; // total sequence length of K or V
- int max_sequence_length; // max sequence length from 4D mask
- int input_hidden_size; // first dimension of weights for input projection
- int hidden_size; // hidden size of Q or K
- int head_size; // hidden size per head of Q or K
- int v_hidden_size; // hidden size of V
- int v_head_size; // hidden size per head of V
+ int kv_sequence_length; // input sequence length of K or V
+ int past_sequence_length; // sequence length in past state of K or V
+ int original_past_sequence_length; // original sequence length in past state of K or V
+ int total_sequence_length; // total sequence length of K or V
+ int max_sequence_length; // max sequence length from 4D mask
+ int input_hidden_size; // first dimension of weights for input projection
+ int hidden_size; // hidden size of Q or K
+ int head_size; // hidden size per head of Q or K
+ int v_hidden_size; // hidden size of V
+ int v_head_size; // hidden size per head of V
int num_heads;
bool is_unidirectional;
bool past_present_share_buffer;
+ bool do_rotary;
float mask_filter_value;
float scale;
AttentionMaskType mask_type;
diff --git a/onnxruntime/contrib_ops/cuda/bert/add_bias_transpose.cu b/onnxruntime/contrib_ops/cuda/bert/add_bias_transpose.cu
index 8f271ecfcb..cd3c33f6c7 100644
--- a/onnxruntime/contrib_ops/cuda/bert/add_bias_transpose.cu
+++ b/onnxruntime/contrib_ops/cuda/bert/add_bias_transpose.cu
@@ -3,32 +3,7 @@
#include "core/providers/cuda/cuda_common.h"
#include "core/providers/cuda/cu_inc/common.cuh"
#include "contrib_ops/cuda/bert/add_bias_transpose.h"
-
-namespace onnxruntime {
-namespace cuda {
-
-struct __align__(8) Half4 {
- half2 x;
- half2 y;
-};
-
-__device__ __forceinline__ Half4 operator+(const Half4& a, const Half4& b) {
- Half4 r;
- r.x = a.x + b.x;
- r.y = a.y + b.y;
- return r;
-}
-
-__device__ __forceinline__ float2 operator+(const float2& a, const float2& b) {
- return make_float2(a.x + b.x, a.y + b.y);
-}
-
-__device__ __forceinline__ float4 operator+(const float4& a, const float4& b) {
- return make_float4(a.x + b.x, a.y + b.y, a.z + b.z, a.w + b.w);
-}
-
-} // namespace cuda
-} // namespace onnxruntime
+#include "contrib_ops/cuda/bert/rotary_embedding_util.h"
using namespace onnxruntime::cuda;
@@ -240,6 +215,151 @@ __global__ void AddBiasTransposeQKV(int M, const T* input, const T* biases, T* o
}
}
+#ifndef USE_ROCM
+template
+__global__ void AddBiasTransposeQKV(int M, const T* input, const T* biases, T* output, T* qkv_add_bias,
+ const int rotary_embedding_dim, const int head_size, const int step,
+ const int format) {
+ // AddBiasTransposeQKV with rotary embedding
+ // Format 1 for unfused attention, or fused causal attention
+ // Input: BxSxMxNxH
+ // Output: MxBxNxSxH
+ // qkv_add_bias: BxSxMxNxH
+ // Format 2 for fused TRT attention
+ // Input: BxSxMxNxH
+ // Output: BxSxNxMxH
+ // qkv_add_bias: BxSxMxNxH
+ // Format 3 for cutlass memory efficient attention
+ // Input: BxSxMxNxH
+ // Output: MxBxSxNxH
+ // B is batch_size, S is sequence_length, M is number of matrices, N is num_heads, H is head_size
+ int n = blockIdx.y;
+ int s = blockIdx.x;
+ int b = blockIdx.z;
+
+ const int seq_len = (gridDim.x == step) ? s : step;
+
+ const int num_heads = gridDim.y;
+
+ const int sequence_length = gridDim.x;
+ const int batch_size = gridDim.z;
+ const int H = head_size;
+ const int NH = num_heads * head_size;
+ const int NHS = NH * sequence_length;
+
+ constexpr int vec_size = Vec_t::size;
+ using Vec_t = typename Vec_t::Type;
+
+ extern __shared__ __align__(sizeof(float2)) char smem_[];
+
+ int tidx = threadIdx.x;
+ const int head_idx = tidx * vec_size;
+
+ if (head_idx < head_size) {
+ const bool is_masked = head_idx >= head_size;
+
+ const int input_offset_base = n * head_size + (s * M) * NH + b * NHS * M;
+ const int src_q_idx = input_offset_base + head_idx;
+ const int src_k_idx = input_offset_base + NH + head_idx;
+ const int src_v_idx = input_offset_base + 2 * NH + head_idx;
+
+ Vec_t q, k, v;
+ Vec_t q_bias, k_bias, v_bias;
+
+ if (!is_masked) {
+ q = *reinterpret_cast(&input[src_q_idx]);
+ k = *reinterpret_cast(&input[src_k_idx]);
+ v = *reinterpret_cast(&input[src_v_idx]);
+
+ q_bias = *reinterpret_cast(&biases[n * H + head_idx]);
+ k_bias = *reinterpret_cast(&biases[NH + n * H + head_idx]);
+ v_bias = *reinterpret_cast(&biases[2 * NH + n * H + head_idx]);
+ }
+
+ q = add(q, q_bias);
+ k = add(k, k_bias);
+ v = add(v, v_bias);
+
+ const bool do_rotary = !is_masked && vec_size * tidx < rotary_embedding_dim;
+
+ T* q_smem = reinterpret_cast(smem_);
+ T* k_smem = q_smem + rotary_embedding_dim;
+
+ const int half_rotary_dim = rotary_embedding_dim / 2;
+ const int half_idx = (head_idx) / half_rotary_dim;
+ const int intra_half_idx = (head_idx) % half_rotary_dim;
+ const int smem_pitch = half_rotary_dim;
+
+ if (do_rotary) {
+ *reinterpret_cast(q_smem + half_idx * smem_pitch + intra_half_idx) = q;
+ *reinterpret_cast(k_smem + half_idx * smem_pitch + intra_half_idx) = k;
+ }
+
+ __syncthreads();
+
+ const int transpose_idx = half_idx * (half_rotary_dim / 2) + intra_half_idx / 2;
+ constexpr int tidx_factor = vec_size / 2;
+
+ if (do_rotary) {
+ vec_from_smem_transpose(q, q_smem, transpose_idx, smem_pitch);
+ vec_from_smem_transpose(k, k_smem, transpose_idx, smem_pitch);
+
+ apply_rotary_embedding(q, k, transpose_idx / tidx_factor, rotary_embedding_dim, seq_len);
+
+ write_smem_transpose(q, q_smem, transpose_idx, smem_pitch);
+ write_smem_transpose(k, k_smem, transpose_idx, smem_pitch);
+ }
+
+ __syncthreads();
+
+ if (do_rotary) {
+ q = *reinterpret_cast(q_smem + half_idx * smem_pitch + intra_half_idx);
+ k = *reinterpret_cast(k_smem + half_idx * smem_pitch + intra_half_idx);
+ }
+
+ int dest_q_idx;
+ int dest_k_idx;
+ int dest_v_idx;
+
+ // Format 1
+ if (format == 1) {
+ const int output_offset_base = s * head_size + n * sequence_length * H + b * NHS;
+ dest_q_idx = output_offset_base + head_idx;
+ dest_k_idx = output_offset_base + NHS * batch_size + head_idx;
+ dest_v_idx = output_offset_base + 2 * NHS * batch_size + head_idx;
+ }
+
+ // Format 2
+ if (format == 2) {
+ const int output_offset_base = M * (b * NHS + s * NH + n * H);
+ dest_q_idx = output_offset_base + head_idx;
+ dest_k_idx = output_offset_base + H + head_idx;
+ dest_v_idx = output_offset_base + 2 * H + head_idx;
+ }
+
+ // Format 3
+ if (format == 3) {
+ const int output_offset_base = n * H + s * NH + b * NHS;
+ dest_q_idx = output_offset_base + head_idx;
+ dest_k_idx = output_offset_base + NHS * batch_size + head_idx;
+ dest_v_idx = output_offset_base + 2 * NHS * batch_size + head_idx;
+ }
+
+ if (!is_masked) {
+ *reinterpret_cast(&output[dest_q_idx]) = q;
+ *reinterpret_cast(&output[dest_k_idx]) = k;
+ *reinterpret_cast(&output[dest_v_idx]) = v;
+
+ if (nullptr != qkv_add_bias) {
+ *reinterpret_cast(&qkv_add_bias[src_q_idx]) = q;
+ *reinterpret_cast(&qkv_add_bias[src_k_idx]) = k;
+ *reinterpret_cast(&qkv_add_bias[src_v_idx]) = v;
+ }
+ }
+ }
+}
+#endif
+
// this suppose 3 matrix in total
template
__global__ void AddBiasTransposeQKV(const T* input, const T* biases, T* output, int v_head_size) {
@@ -320,6 +440,7 @@ __global__ void AddBiasTransposeQKVLarge(const int head_size, const T* input, co
}
}
+
template
__global__ void AddBiasTransposeCutlass(const T* input, const T* biases, T* output, int v_head_size) {
// Format 3 for cutlass memory efficient attention
@@ -518,8 +639,35 @@ template
void InvokeAddBiasTranspose(
cudaStream_t stream, const int num_matrices, const int format, const int max_threads_per_block,
const int batch_size, const int sequence_length, const int num_heads, const int qk_head_size,
- const T* input, const T* biases, T* output, T* qkv_add_bias, const int v_head_size, int total_matrix_count) {
+ const T* input, const T* biases, T* output, T* qkv_add_bias, const int v_head_size, int total_matrix_count,
+ bool do_rotary = false, int original_past_sequence_length = 0) {
assert(num_heads <= max_threads_per_block);
+
+ if (do_rotary) {
+#ifdef USE_ROCM
+ ORT_THROW("Rotary Attention is not supported on ROCm");
+#elif !defined(__CUDA_ARCH__) || __CUDA_ARCH__ >= 530
+ if (format != 1 && format != 2 && format != 3) {
+ ORT_THROW("format must be 1, 2 or 3 for rotary attention");
+ }
+ if (v_head_size != -1 && qk_head_size != v_head_size) {
+ ORT_THROW("qk_head_size must be equal to v_head_size for rotary attention");
+ }
+
+ const int step = original_past_sequence_length == 0 ? sequence_length : original_past_sequence_length;
+ size_t smem_size = 2 * qk_head_size * sizeof(T);
+
+ const dim3 grid(sequence_length, num_heads, batch_size);
+ const dim3 block((qk_head_size / 2 + 31) / 32 * 32, 1, 1);
+ AddBiasTransposeQKV<<>>(total_matrix_count, input, biases, output,
+ qkv_add_bias, qk_head_size, qk_head_size,
+ step, format);
+#else
+ ORT_THROW("Rotary Attention is supported on sm >= 530. Current sm is", __CUDA_ARCH__);
+#endif
+ return;
+ }
+
const dim3 grid(sequence_length, batch_size, num_matrices);
if (qk_head_size * num_heads <= max_threads_per_block) {
const dim3 block(qk_head_size, num_heads, 1);
@@ -575,10 +723,10 @@ template <>
void LaunchAddBiasTranspose(
cudaStream_t stream, const int num_matrices, const int format, const int max_threads_per_block,
const int batch_size, const int sequence_length, const int num_heads, const int qk_head_size,
- const half* input, const half* biases, half* output,
- bool enable_half4, const int v_head_size, half* qkv_add_bias, int total_matrix_count) {
+ const half* input, const half* biases, half* output, bool enable_half4, const int v_head_size,
+ half* qkv_add_bias, int total_matrix_count, bool do_rotary, int original_past_sequence_length) {
total_matrix_count = std::max(num_matrices, total_matrix_count);
- if (enable_half4 && 0 == (qk_head_size % 4) && (v_head_size == -1 || 0 == (v_head_size % 4))) {
+ if (enable_half4 && 0 == (qk_head_size % 4) && (v_head_size == -1 || 0 == (v_head_size % 4)) && !do_rotary) {
const int H = qk_head_size / 4;
const int H_v = v_head_size / 4;
const Half4* input2 = reinterpret_cast(input);
@@ -588,7 +736,7 @@ void LaunchAddBiasTranspose(
InvokeAddBiasTranspose(stream, num_matrices, format, max_threads_per_block,
batch_size, sequence_length, num_heads, H, input2, biases2, output2,
qkv_add_bias2, H_v, total_matrix_count);
- } else if (0 == (qk_head_size & 1) && (v_head_size == -1 || 0 == (v_head_size & 1))) {
+ } else if (0 == (qk_head_size & 1) && (v_head_size == -1 || 0 == (v_head_size & 1)) && !do_rotary) {
const int H = qk_head_size / 2;
const int H_v = v_head_size / 2;
const half2* input2 = reinterpret_cast(input);
@@ -602,7 +750,7 @@ void LaunchAddBiasTranspose(
InvokeAddBiasTranspose(
stream, num_matrices, format, max_threads_per_block,
batch_size, sequence_length, num_heads, qk_head_size, input, biases, output,
- qkv_add_bias, v_head_size, total_matrix_count);
+ qkv_add_bias, v_head_size, total_matrix_count, do_rotary, original_past_sequence_length);
}
}
@@ -610,10 +758,11 @@ template <>
void LaunchAddBiasTranspose(
cudaStream_t stream, const int num_matrices, const int format, const int max_threads_per_block,
const int batch_size, const int sequence_length, const int num_heads, const int qk_head_size,
- const float* input, const float* biases, float* output,
- bool /*enable_half4*/, const int v_head_size, float* qkv_add_bias, int total_matrix_count) {
+ const float* input, const float* biases, float* output, bool /*enable_half4*/,
+ const int v_head_size, float* qkv_add_bias, int total_matrix_count, bool do_rotary,
+ int original_past_sequence_length) {
total_matrix_count = std::max(num_matrices, total_matrix_count);
- if (0 == (qk_head_size % 4) && (v_head_size == -1 || 0 == (v_head_size % 4))) {
+ if (0 == (qk_head_size % 4) && (v_head_size == -1 || 0 == (v_head_size % 4)) && !do_rotary) {
const int H = qk_head_size / 4;
const float4* input2 = reinterpret_cast(input);
const float4* biases2 = reinterpret_cast(biases);
@@ -623,7 +772,7 @@ void LaunchAddBiasTranspose(
stream, num_matrices, format, max_threads_per_block,
batch_size, sequence_length, num_heads, H, input2, biases2, output2,
qkv_add_bias2, v_head_size / 4, total_matrix_count);
- } else if (0 == (qk_head_size & 1) && (v_head_size == -1 || 0 == (v_head_size & 1))) {
+ } else if (0 == (qk_head_size & 1) && (v_head_size == -1 || 0 == (v_head_size & 1)) && !do_rotary) {
const int H = qk_head_size / 2;
const float2* input2 = reinterpret_cast(input);
const float2* biases2 = reinterpret_cast(biases);
@@ -637,7 +786,7 @@ void LaunchAddBiasTranspose(
InvokeAddBiasTranspose(
stream, num_matrices, format, max_threads_per_block,
batch_size, sequence_length, num_heads, qk_head_size, input, biases, output,
- qkv_add_bias, v_head_size, total_matrix_count);
+ qkv_add_bias, v_head_size, total_matrix_count, do_rotary, original_past_sequence_length);
}
}
diff --git a/onnxruntime/contrib_ops/cuda/bert/add_bias_transpose.h b/onnxruntime/contrib_ops/cuda/bert/add_bias_transpose.h
index a2c3265284..86e1983b56 100644
--- a/onnxruntime/contrib_ops/cuda/bert/add_bias_transpose.h
+++ b/onnxruntime/contrib_ops/cuda/bert/add_bias_transpose.h
@@ -33,7 +33,7 @@ void LaunchAddBiasTranspose(
cudaStream_t stream, const int num_matrices, const int format, const int max_threads_per_block,
const int batch_size, const int sequence_length, const int num_heads, const int qk_head_size,
const T* input, const T* biases, T* output, bool enable_half4, const int v_head_size, T* qkv_add_bias = nullptr,
- int total_matrix_count = -1);
+ int total_matrix_count = -1, bool do_rotary = false, int original_past_sequence_length = 0);
// Add (bias) and Transpose for separated inputs of Q, K and V, and output Trt format.
// For self attention:
diff --git a/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu b/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu
index c8ff075d24..214c7dfefd 100644
--- a/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu
+++ b/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu
@@ -330,8 +330,8 @@ Status PrepareQkv(contrib::AttentionParameters& parameters,
// format 2: BxSx(NH + NH + NH) => BxSxNx(H + H + H)
LaunchAddBiasTranspose(stream, matrix_to_transpose, format, max_threads_per_block,
batch_size, sequence_length, num_heads, qk_head_size,
- data.gemm_buffer, data.bias, qkv,
- true, v_head_size, qkv_add_bias, 3);
+ data.gemm_buffer, data.bias, qkv, true, v_head_size, qkv_add_bias,
+ 3, parameters.do_rotary, parameters.original_past_sequence_length);
}
} else if (data.key == nullptr) { // gemm_buffer == nullptr and packed qkv
assert(data.bias == nullptr);
diff --git a/onnxruntime/contrib_ops/cuda/bert/rotary_embedding_util.h b/onnxruntime/contrib_ops/cuda/bert/rotary_embedding_util.h
new file mode 100644
index 0000000000..c232b0e3ce
--- /dev/null
+++ b/onnxruntime/contrib_ops/cuda/bert/rotary_embedding_util.h
@@ -0,0 +1,1292 @@
+/*
+ * Copyright (c) 2020-2022, NVIDIA CORPORATION. All rights reserved.
+ *
+ * Licensed under the Apache License, Version 2.0 (the "License");
+ * you may not use this file except in compliance with the License.
+ * You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+#pragma once
+
+#include "core/providers/cuda/cuda_common.h"
+#include "core/providers/cuda/cu_inc/common.cuh"
+
+namespace onnxruntime {
+namespace cuda {
+
+struct __align__(8) Half4 {
+ half2 x;
+ half2 y;
+};
+
+__device__ __forceinline__ Half4 operator+(const Half4& a, const Half4& b) {
+ Half4 r;
+ r.x = a.x + b.x;
+ r.y = a.y + b.y;
+ return r;
+}
+
+__device__ __forceinline__ float2 operator+(const float2& a, const float2& b) {
+ return make_float2(a.x + b.x, a.y + b.y);
+}
+
+__device__ __forceinline__ float4 operator+(const float4& a, const float4& b) {
+ return make_float4(a.x + b.x, a.y + b.y, a.z + b.z, a.w + b.w);
+}
+
+#ifndef USE_ROCM
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+struct Float8_ {
+ float2 x;
+ float2 y;
+ float2 z;
+ float2 w;
+};
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+struct Float4_ {
+ float2 x;
+ float2 y;
+};
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template
+struct num_elems;
+template<>
+struct num_elems {
+ static constexpr int value = 1;
+};
+template<>
+struct num_elems {
+ static constexpr int value = 2;
+};
+template<>
+struct num_elems {
+ static constexpr int value = 4;
+};
+template<>
+struct num_elems {
+ static constexpr int value = 4;
+};
+template<>
+struct num_elems {
+ static constexpr int value = 8;
+};
+
+template<>
+struct num_elems {
+ static constexpr int value = 2;
+};
+template<>
+struct num_elems {
+ static constexpr int value = 4;
+};
+template<>
+struct num_elems {
+ static constexpr int value = 8;
+};
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template
+struct Vec_t {
+ static constexpr int size = 0;
+};
+
+template<>
+struct Vec_t {
+ using Type = float2;
+ static constexpr int size = 2;
+};
+
+template<>
+struct Vec_t {
+ using Type = float4;
+ static constexpr int size = 4;
+};
+
+template<>
+struct Vec_t {
+ using Type = Float8_;
+ static constexpr int size = 8;
+};
+
+template<>
+struct Vec_t {
+ using Type = uint32_t;
+ static constexpr int size = 2;
+};
+
+template<>
+struct Vec_t {
+ using Type = uint2;
+ static constexpr int size = 4;
+};
+
+template<>
+struct Vec_t {
+ using Type = uint4;
+ static constexpr int size = 8;
+};
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float add(float a, float b)
+{
+ return a + b;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float2 add(float2 a, float2 b)
+{
+ float2 c;
+ c.x = add(a.x, b.x);
+ c.y = add(a.y, b.y);
+ return c;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float4 add(float4 a, float4 b)
+{
+ float4 c;
+ c.x = add(a.x, b.x);
+ c.y = add(a.y, b.y);
+ c.z = add(a.z, b.z);
+ c.w = add(a.w, b.w);
+ return c;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ Float8_ add(Float8_ a, Float8_ b)
+{
+ return;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ uint16_t add(uint16_t a, uint16_t b)
+{
+ uint16_t c;
+ asm volatile("add.f16 %0, %1, %2;\n" : "=h"(c) : "h"(a), "h"(b));
+ return c;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ uint32_t add(uint32_t a, uint32_t b)
+{
+ uint32_t c;
+ asm volatile("add.f16x2 %0, %1, %2;\n" : "=r"(c) : "r"(a), "r"(b));
+ return c;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ uint2 add(uint2 a, uint2 b)
+{
+ uint2 c;
+ c.x = add(a.x, b.x);
+ c.y = add(a.y, b.y);
+ return c;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ uint4 add(uint4 a, uint4 b)
+{
+ uint4 c;
+ c.x = add(a.x, b.x);
+ c.y = add(a.y, b.y);
+ c.z = add(a.z, b.z);
+ c.w = add(a.w, b.w);
+ return c;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ uint16_t float_to_half(float f)
+{
+ union {
+ uint32_t u32;
+ uint16_t u16[2];
+ } tmp;
+#if 0 && defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800 // Is it better?
+ float zero = 0.f;
+ asm volatile("cvt.rn.f16x2.f32 %0, %1, %2;\n" : "=r"(tmp.u32) : "f"(zero), "f"(f));
+#else
+ asm volatile("cvt.rn.f16.f32 %0, %1;\n" : "=h"(tmp.u16[0]) : "f"(f));
+#endif
+ return tmp.u16[0];
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ uint32_t float2_to_half2(float2 f)
+{
+ union {
+ uint32_t u32;
+ uint16_t u16[2];
+ } tmp;
+#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800
+ asm volatile("cvt.rn.f16x2.f32 %0, %1, %2;\n" : "=r"(tmp.u32) : "f"(f.y), "f"(f.x));
+#else
+ asm volatile("cvt.rn.f16.f32 %0, %1;\n" : "=h"(tmp.u16[0]) : "f"(f.x));
+ asm volatile("cvt.rn.f16.f32 %0, %1;\n" : "=h"(tmp.u16[1]) : "f"(f.y));
+#endif
+ return tmp.u32;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float half_to_float(uint16_t h)
+{
+ float f;
+ asm volatile("cvt.f32.f16 %0, %1;\n" : "=f"(f) : "h"(h));
+ return f;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float2 half2_to_float2(uint32_t v)
+{
+ uint16_t lo, hi;
+ asm volatile("mov.b32 {%0, %1}, %2;\n" : "=h"(lo), "=h"(hi) : "r"(v));
+ return make_float2(half_to_float(lo), half_to_float(hi));
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float add(float a, uint16_t b)
+{
+ return a + half_to_float(b);
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float2 add(uint32_t a, float2 fb)
+{
+ float2 fa = half2_to_float2(a);
+ return add(fa, fb);
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ Float4_ add(uint2 a, Float4_ fb)
+{
+ Float4_ fc;
+ fc.x = add(a.x, fb.x);
+ fc.y = add(a.y, fb.y);
+ return fc;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ Float8_ add(uint4 a, Float8_ fb)
+{
+ Float8_ fc;
+ fc.x = add(a.x, fb.x);
+ fc.y = add(a.y, fb.y);
+ fc.z = add(a.z, fb.z);
+ fc.w = add(a.w, fb.w);
+ return fc;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ uint32_t h0_h0(uint16_t a)
+{
+ uint32_t b;
+ asm volatile("mov.b32 %0, {%1, %1};" : "=r"(b) : "h"(a));
+ return b;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float fma(float a, float b, float c)
+{
+ return a * b + c;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float2 fma(float2 a, float2 b, float2 c)
+{
+ float2 d;
+ d.x = fma(a.x, b.x, c.x);
+ d.y = fma(a.y, b.y, c.y);
+ return d;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float2 fma(float a, float2 b, float2 c)
+{
+ float2 d;
+ d.x = fma(a, b.x, c.x);
+ d.y = fma(a, b.y, c.y);
+ return d;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float4 fma(float4 a, float4 b, float4 c)
+{
+ float4 d;
+ d.x = fma(a.x, b.x, c.x);
+ d.y = fma(a.y, b.y, c.y);
+ d.z = fma(a.z, b.z, c.z);
+ d.w = fma(a.w, b.w, c.w);
+ return d;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float4 fma(float a, float4 b, float4 c)
+{
+ float4 d;
+ d.x = fma(a, b.x, c.x);
+ d.y = fma(a, b.y, c.y);
+ d.z = fma(a, b.z, c.z);
+ d.w = fma(a, b.w, c.w);
+ return d;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ Float4_ fma(float a, Float4_ b, Float4_ c)
+{
+ Float4_ d;
+ d.x = fma(a, b.x, c.x);
+ d.y = fma(a, b.y, c.y);
+ return d;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ Float8_ fma(float a, Float8_ b, Float8_ c)
+{
+ Float8_ d;
+ d.x = fma(a, b.x, c.x);
+ d.y = fma(a, b.y, c.y);
+ d.z = fma(a, b.z, c.z);
+ d.w = fma(a, b.w, c.w);
+ return d;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ uint32_t fma(uint32_t a, uint32_t b, uint32_t c)
+{
+ uint32_t d;
+ asm volatile("fma.rn.f16x2 %0, %1, %2, %3;\n" : "=r"(d) : "r"(a), "r"(b), "r"(c));
+ return d;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ uint32_t fma(uint16_t a, uint32_t b, uint32_t c)
+{
+ return fma(h0_h0(a), b, c);
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ uint2 fma(uint2 a, uint2 b, uint2 c)
+{
+ uint2 d;
+ d.x = fma(a.x, b.x, c.x);
+ d.y = fma(a.y, b.y, c.y);
+ return d;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ uint2 fma(uint16_t a, uint2 b, uint2 c)
+{
+ uint32_t s = h0_h0(a);
+ uint2 d;
+ d.x = fma(s, b.x, c.x);
+ d.y = fma(s, b.y, c.y);
+ return d;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ uint4 fma(uint4 a, uint4 b, uint4 c)
+{
+ uint4 d;
+ d.x = fma(a.x, b.x, c.x);
+ d.y = fma(a.y, b.y, c.y);
+ d.z = fma(a.z, b.z, c.z);
+ d.w = fma(a.w, b.w, c.w);
+ return d;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ uint4 fma(uint16_t a, uint4 b, uint4 c)
+{
+ uint32_t s = h0_h0(a);
+ uint4 d;
+ d.x = fma(s, b.x, c.x);
+ d.y = fma(s, b.y, c.y);
+ d.z = fma(s, b.z, c.z);
+ d.w = fma(s, b.w, c.w);
+ return d;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float fma(uint16_t a, uint16_t b, float fc)
+{
+ float fa = half_to_float(a);
+ float fb = half_to_float(b);
+ return fa * fb + fc;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float2 fma(uint32_t a, uint32_t b, float2 fc)
+{
+ float2 fa = half2_to_float2(a);
+ float2 fb = half2_to_float2(b);
+ return fma(fa, fb, fc);
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float2 fma(uint16_t a, uint32_t b, float2 fc)
+{
+ return fma(h0_h0(a), b, fc);
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ Float4_ fma(uint2 a, uint2 b, Float4_ fc)
+{
+ Float4_ fd;
+ fd.x = fma(a.x, b.x, fc.x);
+ fd.y = fma(a.y, b.y, fc.y);
+ return fd;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ Float4_ fma(uint16_t a, uint2 b, Float4_ fc)
+{
+ uint32_t s = h0_h0(a);
+ Float4_ fd;
+ fd.x = fma(s, b.x, fc.x);
+ fd.y = fma(s, b.y, fc.y);
+ return fd;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ Float8_ fma(uint4 a, uint4 b, Float8_ fc)
+{
+ Float8_ fd;
+ fd.x = fma(a.x, b.x, fc.x);
+ fd.y = fma(a.y, b.y, fc.y);
+ fd.z = fma(a.z, b.z, fc.z);
+ fd.w = fma(a.w, b.w, fc.w);
+ return fd;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ Float8_ fma(uint16_t a, uint4 b, Float8_ fc)
+{
+ uint32_t s = h0_h0(a);
+ Float8_ fd;
+ fd.x = fma(s, b.x, fc.x);
+ fd.y = fma(s, b.y, fc.y);
+ fd.z = fma(s, b.z, fc.z);
+ fd.w = fma(s, b.w, fc.w);
+ return fd;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template
+inline __device__ Acc mul(A a, B b)
+{
+ return a * b;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ float mul(float a, float b)
+{
+ return a * b;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ float2 mul(float2 a, float2 b)
+{
+ float2 c;
+ c.x = a.x * b.x;
+ c.y = a.y * b.y;
+ return c;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ float2 mul(float a, float2 b)
+{
+ float2 c;
+ c.x = a * b.x;
+ c.y = a * b.y;
+ return c;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ float4 mul(float4 a, float4 b)
+{
+ float4 c;
+ c.x = a.x * b.x;
+ c.y = a.y * b.y;
+ c.z = a.z * b.z;
+ c.w = a.w * b.w;
+ return c;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ float4 mul(float a, float4 b)
+{
+ float4 c;
+ c.x = a * b.x;
+ c.y = a * b.y;
+ c.z = a * b.z;
+ c.w = a * b.w;
+ return c;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ Float8_ mul(float a, Float8_ b)
+{
+ Float8_ c;
+ c.x = make_float2(a * b.x.x, a * b.x.y);
+ c.y = make_float2(a * b.y.x, a * b.y.y);
+ c.z = make_float2(a * b.z.x, a * b.z.y);
+ c.w = make_float2(a * b.w.x, a * b.w.y);
+ return c;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ uint16_t mul(uint16_t a, uint16_t b)
+{
+ uint16_t c;
+ asm volatile("mul.f16 %0, %1, %2;\n" : "=h"(c) : "h"(a), "h"(b));
+ return c;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ uint32_t mul(uint32_t a, uint32_t b)
+{
+ uint32_t c;
+ asm volatile("mul.f16x2 %0, %1, %2;\n" : "=r"(c) : "r"(a), "r"(b));
+ return c;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ uint32_t mul(uint16_t a, uint32_t b)
+{
+ return mul(h0_h0(a), b);
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ uint2 mul(uint2 a, uint2 b)
+{
+ uint2 c;
+ c.x = mul(a.x, b.x);
+ c.y = mul(a.y, b.y);
+ return c;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ uint2 mul(uint16_t a, uint2 b)
+{
+ uint32_t s = h0_h0(a);
+ uint2 c;
+ c.x = mul(s, b.x);
+ c.y = mul(s, b.y);
+ return c;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ uint4 mul(uint4 a, uint4 b)
+{
+ uint4 c;
+ c.x = mul(a.x, b.x);
+ c.y = mul(a.y, b.y);
+ c.z = mul(a.z, b.z);
+ c.w = mul(a.w, b.w);
+ return c;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ uint4 mul(uint16_t a, uint4 b)
+{
+ uint32_t s = h0_h0(a);
+ uint4 c;
+ c.x = mul(s, b.x);
+ c.y = mul(s, b.y);
+ c.z = mul(s, b.z);
+ c.w = mul(s, b.w);
+ return c;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ float mul(uint16_t a, uint16_t b)
+{
+ float fa = half_to_float(a);
+ float fb = half_to_float(b);
+ return fa * fb;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ float mul(uint16_t a, float b)
+{
+ return half_to_float(a) * b;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ float2 mul(uint32_t a, uint32_t b)
+{
+ float2 fa = half2_to_float2(a);
+ float2 fb = half2_to_float2(b);
+ return mul(fa, fb);
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ float2 mul(uint16_t a, uint32_t b)
+{
+ return mul(h0_h0(a), b);
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ Float4_ mul(uint2 a, uint2 b)
+{
+ Float4_ fc;
+ fc.x = mul(a.x, b.x);
+ fc.y = mul(a.y, b.y);
+ return fc;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ Float4_ mul(uint16_t a, uint2 b)
+{
+ uint32_t s = h0_h0(a);
+ Float4_ fc;
+ fc.x = mul(s, b.x);
+ fc.y = mul(s, b.y);
+ return fc;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ Float8_ mul(uint4 a, uint4 b)
+{
+ Float8_ fc;
+ fc.x = mul(a.x, b.x);
+ fc.y = mul(a.y, b.y);
+ fc.z = mul(a.z, b.z);
+ fc.w = mul(a.w, b.w);
+ return fc;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template<>
+inline __device__ Float8_ mul(uint16_t a, uint4 b)
+{
+ uint32_t s = h0_h0(a);
+ Float8_ fc;
+ fc.x = mul(s, b.x);
+ fc.y = mul(s, b.y);
+ fc.z = mul(s, b.z);
+ fc.w = mul(s, b.w);
+ return fc;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float sum(float v)
+{
+ return v;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float sum(float2 v)
+{
+ return v.x + v.y;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float sum(float4 v)
+{
+ return v.x + v.y + v.z + v.w;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float sum(uint16_t v)
+{
+ return half_to_float(v);
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float sum(uint32_t v)
+{
+ float2 tmp = half2_to_float2(v);
+ return tmp.x + tmp.y;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float sum(uint2 v)
+{
+ uint32_t c = add(v.x, v.y);
+ return sum(c);
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float sum(uint4 v)
+{
+#if 1
+ uint32_t c = add(v.x, v.y);
+ c = add(c, v.z);
+ c = add(c, v.w);
+#else
+ uint32_t c = add(v.x, v.y);
+ uint32_t d = add(v.z, v.w);
+ c = add(c, d);
+#endif
+ return sum(c);
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float sum(Float4_ v)
+{
+ return v.x.x + v.x.y + v.y.x + v.y.y;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float sum(Float8_ v)
+{
+ return v.x.x + v.x.y + v.y.x + v.y.y + v.z.x + v.z.y + v.w.x + v.w.y;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template
+inline __device__ float dot(T a, T b)
+{
+ return sum(mul(a, b));
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template
+inline __device__ float dot(T a, T b)
+{
+ return sum(mul(a, b));
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ void zero(uint16_t& dst)
+{
+ dst = uint16_t(0);
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+template
+inline __device__ void zero(T& dst)
+{
+ constexpr int WORDS = sizeof(T) / 4;
+ union {
+ T raw;
+ uint32_t words[WORDS];
+ } tmp;
+#pragma unroll
+ for (int ii = 0; ii < WORDS; ++ii) {
+ tmp.words[ii] = 0u;
+ }
+ dst = tmp.raw;
+}
+
+////////////////////////////////////////////////////////////////////////////////////////////////////
+
+inline __device__ float2 rotary_embedding_coefficient(const int zid, const int rot_embed_dim, const float t_step)
+{
+ const float inv_freq = t_step / pow(10000.0f, zid / (float)rot_embed_dim);
+ return {cos(inv_freq), sin(inv_freq)};
+}
+
+inline __device__ float2 rotary_embedding_transform(const float2 v, const float2 coef)
+{
+ float2 rot_v;
+ rot_v.x = coef.x * v.x - coef.y * v.y;
+ rot_v.y = coef.x * v.y + coef.y * v.x;
+ return rot_v;
+}
+
+inline __device__ uint32_t rotary_embedding_transform(const uint32_t v, const float2 coef)
+{
+ float2 fv = half2_to_float2(v);
+ float2 rot_fv = rotary_embedding_transform(fv, coef);
+ return float2_to_half2(rot_fv);
+}
+
+inline __device__ void apply_rotary_embedding(float& q, int zid, int rot_embed_dim, int t_step)
+{
+ return;
+}
+
+inline __device__ void apply_rotary_embedding(float& q, float& k, int zid, int rot_embed_dim, int t_step)
+{
+ return;
+}
+
+inline __device__ void apply_rotary_embedding(Float8_& q, Float8_& k, int zid, int rot_embed_dim, int t_step)
+{
+ return;
+}
+
+inline __device__ void apply_rotary_embedding(float2& q, int tid, int rot_embed_dim, int t_step)
+{
+ if (2 * tid >= rot_embed_dim) {
+ return;
+ }
+ const auto coef = rotary_embedding_coefficient(2 * tid, rot_embed_dim, t_step);
+ q = rotary_embedding_transform(q, coef);
+}
+
+inline __device__ void apply_rotary_embedding(float2& q, float2& k, int tid, int rot_embed_dim, int t_step)
+{
+ if (2 * tid >= rot_embed_dim) {
+ return;
+ }
+ const auto coef = rotary_embedding_coefficient(2 * tid, rot_embed_dim, t_step);
+ q = rotary_embedding_transform(q, coef);
+ k = rotary_embedding_transform(k, coef);
+}
+
+inline __device__ void apply_rotary_embedding(float4& q, int tid, int rot_embed_dim, int t_step)
+{
+ if (4 * tid >= rot_embed_dim) {
+ return;
+ }
+
+ Float4_& q_ = *reinterpret_cast(&q);
+ const auto coef0 = rotary_embedding_coefficient(4 * tid, rot_embed_dim, t_step);
+ q_.x = rotary_embedding_transform(q_.x, coef0);
+ const auto coef1 = rotary_embedding_coefficient(4 * tid + 2, rot_embed_dim, t_step);
+ q_.y = rotary_embedding_transform(q_.y, coef1);
+}
+
+inline __device__ void apply_rotary_embedding(float4& q, float4& k, int tid, int rot_embed_dim, int t_step)
+{
+ if (4 * tid >= rot_embed_dim) {
+ return;
+ }
+
+ Float4_& q_ = *reinterpret_cast(&q);
+ Float4_& k_ = *reinterpret_cast(&k);
+ const auto coef0 = rotary_embedding_coefficient(4 * tid, rot_embed_dim, t_step);
+ q_.x = rotary_embedding_transform(q_.x, coef0);
+ k_.x = rotary_embedding_transform(k_.x, coef0);
+ const auto coef1 = rotary_embedding_coefficient(4 * tid + 2, rot_embed_dim, t_step);
+ q_.y = rotary_embedding_transform(q_.y, coef1);
+ k_.y = rotary_embedding_transform(k_.y, coef1);
+}
+
+inline __device__ void apply_rotary_embedding(uint32_t& q, int tid, int rot_embed_dim, int t_step)
+{
+ if (2 * tid >= rot_embed_dim) {
+ return;
+ }
+ const auto coef = rotary_embedding_coefficient(2 * tid, rot_embed_dim, t_step);
+ q = rotary_embedding_transform(q, coef);
+}
+
+inline __device__ void apply_rotary_embedding(uint32_t& q, uint32_t& k, int tid, int rot_embed_dim, int t_step)
+{
+ if (2 * tid >= rot_embed_dim) {
+ return;
+ }
+ const auto coef = rotary_embedding_coefficient(2 * tid, rot_embed_dim, t_step);
+ q = rotary_embedding_transform(q, coef);
+ k = rotary_embedding_transform(k, coef);
+}
+
+inline __device__ void apply_rotary_embedding(uint2& q, int tid, int rot_embed_dim, int t_step)
+{
+ if (4 * tid >= rot_embed_dim) {
+ return;
+ }
+ const auto coef0 = rotary_embedding_coefficient(4 * tid, rot_embed_dim, t_step);
+ q.x = rotary_embedding_transform(q.x, coef0);
+ const auto coef1 = rotary_embedding_coefficient(4 * tid + 2, rot_embed_dim, t_step);
+ q.y = rotary_embedding_transform(q.y, coef1);
+}
+
+inline __device__ void apply_rotary_embedding(uint2& q, uint2& k, int tid, int rot_embed_dim, int t_step)
+{
+ if (4 * tid >= rot_embed_dim) {
+ return;
+ }
+ const auto coef0 = rotary_embedding_coefficient(4 * tid, rot_embed_dim, t_step);
+ q.x = rotary_embedding_transform(q.x, coef0);
+ k.x = rotary_embedding_transform(k.x, coef0);
+ const auto coef1 = rotary_embedding_coefficient(4 * tid + 2, rot_embed_dim, t_step);
+ q.y = rotary_embedding_transform(q.y, coef1);
+ k.y = rotary_embedding_transform(k.y, coef1);
+}
+
+inline __device__ void apply_rotary_embedding(uint4& q, int tid, int rot_embed_dim, int t_step)
+{
+ if (8 * tid >= rot_embed_dim) {
+ return;
+ }
+ const auto coef0 = rotary_embedding_coefficient(8 * tid, rot_embed_dim, t_step);
+ q.x = rotary_embedding_transform(q.x, coef0);
+ const auto coef1 = rotary_embedding_coefficient(8 * tid + 2, rot_embed_dim, t_step);
+ q.y = rotary_embedding_transform(q.y, coef1);
+ const auto coef2 = rotary_embedding_coefficient(8 * tid + 4, rot_embed_dim, t_step);
+ q.z = rotary_embedding_transform(q.z, coef2);
+ const auto coef3 = rotary_embedding_coefficient(8 * tid + 6, rot_embed_dim, t_step);
+ q.w = rotary_embedding_transform(q.w, coef3);
+}
+
+inline __device__ void apply_rotary_embedding(uint4& q, uint4& k, int tid, int rot_embed_dim, int t_step)
+{
+ if (8 * tid >= rot_embed_dim) {
+ return;
+ }
+ const auto coef0 = rotary_embedding_coefficient(8 * tid, rot_embed_dim, t_step);
+ q.x = rotary_embedding_transform(q.x, coef0);
+ k.x = rotary_embedding_transform(k.x, coef0);
+ const auto coef1 = rotary_embedding_coefficient(8 * tid + 2, rot_embed_dim, t_step);
+ q.y = rotary_embedding_transform(q.y, coef1);
+ k.y = rotary_embedding_transform(k.y, coef1);
+ const auto coef2 = rotary_embedding_coefficient(8 * tid + 4, rot_embed_dim, t_step);
+ q.z = rotary_embedding_transform(q.z, coef2);
+ k.z = rotary_embedding_transform(k.z, coef2);
+ const auto coef3 = rotary_embedding_coefficient(8 * tid + 6, rot_embed_dim, t_step);
+ q.w = rotary_embedding_transform(q.w, coef3);
+ k.w = rotary_embedding_transform(k.w, coef3);
+}
+
+template
+__device__ __inline__ void vec_from_smem_transpose(Vec_T& vec, T* smem, int transpose_idx, int smem_pitch);
+
+template<>
+__device__ __inline__ void vec_from_smem_transpose(float& vec, float* smem, int transpose_idx, int smem_pitch)
+{
+ return;
+}
+
+template<>
+__device__ __inline__ void vec_from_smem_transpose(float4& vec, float2* smem, int transpose_idx, int smem_pitch)
+{
+ return;
+}
+
+template<>
+__device__ __inline__ void vec_from_smem_transpose(Float8_& vec, float4* smem, int transpose_idx, int smem_pitch)
+{
+ return;
+}
+
+template<>
+__device__ __inline__ void vec_from_smem_transpose(uint2& vec, half2* smem, int transpose_idx, int smem_pitch)
+{
+ return;
+}
+
+template<>
+__device__ __inline__ void vec_from_smem_transpose(uint4& vec, Half4* smem, int transpose_idx, int smem_pitch)
+{
+ return;
+}
+
+template<>
+__device__ __inline__ void vec_from_smem_transpose(uint32_t& vec, uint16_t* smem, int transpose_idx, int smem_pitch)
+{
+ union {
+ uint32_t u32;
+ uint16_t u16[2];
+ } tmp;
+ tmp.u16[0] = smem[transpose_idx];
+ tmp.u16[1] = smem[smem_pitch + transpose_idx];
+
+ vec = tmp.u32;
+}
+
+template<>
+__device__ __inline__ void vec_from_smem_transpose(uint2& vec, uint16_t* smem, int transpose_idx, int smem_pitch)
+{
+ union {
+ uint32_t u32;
+ uint16_t u16[2];
+ } tmp_1, tmp_2;
+ tmp_1.u32 = *reinterpret_cast(&smem[transpose_idx]);
+ tmp_2.u32 = *reinterpret_cast(&smem[smem_pitch + transpose_idx]);
+
+ union {
+ uint2 u32x2;
+ uint16_t u16[4];
+ } tmp_3;
+ tmp_3.u16[0] = tmp_1.u16[0];
+ tmp_3.u16[1] = tmp_2.u16[0];
+ tmp_3.u16[2] = tmp_1.u16[1];
+ tmp_3.u16[3] = tmp_2.u16[1];
+
+ vec = tmp_3.u32x2;
+}
+
+template<>
+__device__ __inline__ void vec_from_smem_transpose(uint4& vec, uint16_t* smem, int transpose_idx, int smem_pitch)
+{
+ union {
+ uint64_t u64;
+ uint16_t u16[4];
+ } tmp_1, tmp_2;
+ tmp_1.u64 = *reinterpret_cast(&smem[transpose_idx]);
+ tmp_2.u64 = *reinterpret_cast(&smem[smem_pitch + transpose_idx]);
+
+ union {
+ uint4 u32x4;
+ uint16_t u16[8];
+ } tmp_3;
+ tmp_3.u16[0] = tmp_1.u16[0];
+ tmp_3.u16[1] = tmp_2.u16[0];
+ tmp_3.u16[2] = tmp_1.u16[1];
+ tmp_3.u16[3] = tmp_2.u16[1];
+ tmp_3.u16[4] = tmp_1.u16[2];
+ tmp_3.u16[5] = tmp_2.u16[2];
+ tmp_3.u16[6] = tmp_1.u16[3];
+ tmp_3.u16[7] = tmp_2.u16[3];
+
+ vec = tmp_3.u32x4;
+}
+
+template<>
+__device__ __inline__ void vec_from_smem_transpose(float4& vec, float* smem, int transpose_idx, int smem_pitch)
+{
+ vec.x = smem[transpose_idx];
+ vec.z = smem[transpose_idx + 1];
+ vec.y = smem[smem_pitch + transpose_idx];
+ vec.w = smem[smem_pitch + transpose_idx + 1];
+}
+
+template<>
+__device__ __inline__ void vec_from_smem_transpose(uint32_t& vec, half* smem, int transpose_idx, int smem_pitch)
+{
+ union {
+ uint32_t u32;
+ half u16[2];
+ } tmp;
+ tmp.u16[0] = smem[transpose_idx];
+ tmp.u16[1] = smem[smem_pitch + transpose_idx];
+
+ vec = tmp.u32;
+}
+
+template<>
+__device__ __inline__ void vec_from_smem_transpose(float2& vec, float* smem, int transpose_idx, int smem_pitch)
+{
+ vec.x = smem[transpose_idx];
+ vec.y = smem[smem_pitch + transpose_idx];
+}
+
+template
+__device__ __inline__ void write_smem_transpose(const Vec_T& vec, T* smem, int transpose_idx, int smem_pitch);
+
+template<>
+__device__ __inline__ void write_smem_transpose(const float& vec, float* smem, int transpose_idx, int smem_pitch)
+{
+ return;
+}
+
+template<>
+__device__ __inline__ void write_smem_transpose(const float4& vec, float2* smem, int transpose_idx, int smem_pitch)
+{
+ return;
+}
+
+template<>
+__device__ __inline__ void write_smem_transpose(const Float8_& vec, float4* smem, int transpose_idx, int smem_pitch)
+{
+ return;
+}
+
+template<>
+__device__ __inline__ void write_smem_transpose(const uint2& vec, half2* smem, int transpose_idx, int smem_pitch)
+{
+ return;
+}
+
+template<>
+__device__ __inline__ void write_smem_transpose(const uint4& vec, Half4* smem, int transpose_idx, int smem_pitch)
+{
+ return;
+}
+
+template<>
+__device__ __inline__ void write_smem_transpose(const uint4& vec, uint16_t* smem, int transpose_idx, int smem_pitch)
+{
+ union {
+ uint64_t u64;
+ uint16_t u16[4];
+ } tmp_1, tmp_2;
+
+ union {
+ uint4 u32x4;
+ uint16_t u16[8];
+ } tmp_3;
+ tmp_3.u32x4 = vec;
+ tmp_1.u16[0] = tmp_3.u16[0];
+ tmp_2.u16[0] = tmp_3.u16[1];
+ tmp_1.u16[1] = tmp_3.u16[2];
+ tmp_2.u16[1] = tmp_3.u16[3];
+ tmp_1.u16[2] = tmp_3.u16[4];
+ tmp_2.u16[2] = tmp_3.u16[5];
+ tmp_1.u16[3] = tmp_3.u16[6];
+ tmp_2.u16[3] = tmp_3.u16[7];
+
+ *reinterpret_cast(&smem[transpose_idx]) = tmp_1.u64;
+ *reinterpret_cast(&smem[smem_pitch + transpose_idx]) = tmp_2.u64;
+}
+
+template<>
+__device__ __inline__ void write_smem_transpose(const uint2& vec, uint16_t* smem, int transpose_idx, int smem_pitch)
+{
+ union {
+ uint32_t u32;
+ uint16_t u16[2];
+ } tmp_1, tmp_2;
+
+ union {
+ uint2 u32x2;
+ uint16_t u16[4];
+ } tmp_3;
+ tmp_3.u32x2 = vec;
+ tmp_1.u16[0] = tmp_3.u16[0];
+ tmp_2.u16[0] = tmp_3.u16[1];
+ tmp_1.u16[1] = tmp_3.u16[2];
+ tmp_2.u16[1] = tmp_3.u16[3];
+
+ *reinterpret_cast(&smem[transpose_idx]) = tmp_1.u32;
+ *reinterpret_cast(&smem[smem_pitch + transpose_idx]) = tmp_2.u32;
+}
+
+template<>
+__device__ __inline__ void write_smem_transpose(const uint32_t& vec, uint16_t* smem, int transpose_idx, int smem_pitch)
+{
+ union {
+ uint32_t u32;
+ uint16_t u16[2];
+ } tmp;
+ tmp.u32 = vec;
+
+ smem[transpose_idx] = tmp.u16[0];
+ smem[smem_pitch + transpose_idx] = tmp.u16[1];
+}
+
+template<>
+__device__ __inline__ void write_smem_transpose(const float4& vec, float* smem, int transpose_idx, int smem_pitch)
+{
+ smem[transpose_idx] = vec.x;
+ smem[transpose_idx + 1] = vec.z;
+ smem[smem_pitch + transpose_idx] = vec.y;
+ smem[smem_pitch + transpose_idx + 1] = vec.w;
+}
+
+template<>
+__device__ __inline__ void write_smem_transpose(const uint32_t& vec, half* smem, int transpose_idx, int smem_pitch)
+{
+ union {
+ uint32_t u32;
+ half u16[2];
+ } tmp;
+
+ tmp.u32 = vec;
+ smem[transpose_idx] = tmp.u16[0];
+ smem[smem_pitch + transpose_idx] = tmp.u16[1];
+}
+
+template<>
+__device__ __inline__ void write_smem_transpose(const float2& vec, float* smem, int transpose_idx, int smem_pitch)
+{
+ smem[transpose_idx] = vec.x;
+ smem[smem_pitch + transpose_idx] = vec.y;
+}
+
+#endif
+
+} // namespace onnxruntime
+} // namespace cuda
diff --git a/onnxruntime/core/graph/contrib_ops/bert_defs.cc b/onnxruntime/core/graph/contrib_ops/bert_defs.cc
index 580a0a9934..e3a2086741 100644
--- a/onnxruntime/core/graph/contrib_ops/bert_defs.cc
+++ b/onnxruntime/core/graph/contrib_ops/bert_defs.cc
@@ -231,6 +231,10 @@ ONNX_MS_OPERATOR_SET_SCHEMA(
"(2, batch_size, num_heads, max_sequence_length, head_size)",
AttributeProto::INT,
OPTIONAL_VALUE)
+ .Attr("do_rotary",
+ "Whether to use rotary position embedding. Default value is 0.",
+ AttributeProto::INT,
+ OPTIONAL_VALUE)
.Attr("mask_filter_value",
"The value to be filled in the attention mask. Default value is -10000.0f",
AttributeProto::FLOAT,
diff --git a/onnxruntime/core/graph/contrib_ops/quantization_defs.cc b/onnxruntime/core/graph/contrib_ops/quantization_defs.cc
index 91e4f5d8ff..375c5878d1 100644
--- a/onnxruntime/core/graph/contrib_ops/quantization_defs.cc
+++ b/onnxruntime/core/graph/contrib_ops/quantization_defs.cc
@@ -949,6 +949,8 @@ ONNX_MS_OPERATOR_SET_SCHEMA(
.Attr("num_heads", "Number of attention heads", AttributeProto::INT)
.Attr("unidirectional", "Whether every token can only attend to previous tokens. Default value is 0.",
AttributeProto::INT, static_cast(0))
+ .Attr("do_rotary", "Whether to use rotary position embedding. Default value is 0.",
+ AttributeProto::INT, OPTIONAL_VALUE)
.Attr("past_present_share_buffer", "Corresponding past and present are same tensor, its shape is "
"(2, batch_size, num_heads, max_sequence_length, head_size)",
AttributeProto::INT, OPTIONAL_VALUE)
diff --git a/onnxruntime/test/contrib_ops/attention_op_test.cc b/onnxruntime/test/contrib_ops/attention_op_test.cc
index daeec7a64c..709fbcb378 100644
--- a/onnxruntime/test/contrib_ops/attention_op_test.cc
+++ b/onnxruntime/test/contrib_ops/attention_op_test.cc
@@ -62,7 +62,8 @@ static void RunAttentionTest(
const std::vector& relative_position_bias_data = {},
int kv_sequence_length = 0,
bool past_present_share_buffer = false,
- bool use_scale = false) {
+ bool use_scale = false,
+ bool do_neox_rotary = false) {
input_hidden_size = (input_hidden_size == 0 ? hidden_size : input_hidden_size); // By default, no pruning.
kv_sequence_length = (kv_sequence_length == 0 ? sequence_length : kv_sequence_length);
past_present_share_buffer = past_present_share_buffer && use_past_state;
@@ -82,6 +83,9 @@ static void RunAttentionTest(
if (use_scale && !enable_rocm) {
tester.AddAttribute("scale", static_cast(1.f / sqrt(head_size)));
}
+ if (do_neox_rotary && !enable_rocm) {
+ tester.AddAttribute("do_rotary", static_cast(do_neox_rotary ? 1 : 0));
+ }
int32_t qkv_hidden_size_sum;
int32_t v_hidden_size;
@@ -267,19 +271,20 @@ static void RunAttentionTest(
const std::vector& relative_position_bias_data = {},
int kv_sequence_length = 0,
bool past_present_share_buffer = false,
- bool use_scale = false) {
+ bool use_scale = false,
+ bool do_neox_rotary = false) {
RunAttentionTest(input_data, weights_data, false, bias_data, mask_index_data, output_data,
batch_size, sequence_length, hidden_size, number_of_heads,
use_float16, is_unidirectional, use_past_state, past_sequence_length,
past_data, present_data, mask_type, input_hidden_size, max_sequence_length,
disable_cpu, disable_cuda, disable_rocm, qkv_sizes, relative_position_bias_data,
- kv_sequence_length, past_present_share_buffer, use_scale);
+ kv_sequence_length, past_present_share_buffer, use_scale, do_neox_rotary);
RunAttentionTest(input_data, weights_data, true, bias_data, mask_index_data, output_data,
batch_size, sequence_length, hidden_size, number_of_heads,
use_float16, is_unidirectional, use_past_state, past_sequence_length,
past_data, present_data, mask_type, input_hidden_size, max_sequence_length,
disable_cpu, disable_cuda, disable_rocm, qkv_sizes, relative_position_bias_data,
- kv_sequence_length, past_present_share_buffer, use_scale);
+ kv_sequence_length, past_present_share_buffer, use_scale, do_neox_rotary);
}
TEST(AttentionTest, AttentionBatch1) {
@@ -1716,6 +1721,51 @@ TEST(AttentionTest, AttentionWithNormFactor) {
true /*use_scale*/);
}
+TEST(AttentionTest, AttentionWithNeoXRotaryEmbedding) {
+ int batch_size = 2;
+ int sequence_length = 2;
+ int hidden_size = 4;
+ int number_of_heads = 2;
+
+ std::vector input_data = {
+ 0.5f, 0.2f, 0.3f, -0.6f,
+ 0.8f, -0.5f, 0.0f, 1.f,
+ 0.8f, -0.5f, 0.0f, 1.f,
+ 0.5f, 0.2f, 0.3f, -0.6f};
+
+ std::vector weight_data = {
+ 0.1f, -0.2f, 0.3f, 1.0f, 1.1f, 0.3f, 0.5f, 0.2f, 0.3f, -0.6f, 1.5f, 2.0f,
+ 0.5f, 0.1f, 0.4f, 1.6f, 1.0f, 2.0f, 0.4f, 0.8f, 0.9f, 0.1f, -1.3f, 0.7f,
+ 0.3f, 0.2f, 4.0f, 2.2f, 1.6f, 1.1f, 0.7f, 0.2f, 0.4f, 1.0f, 1.2f, 0.5f,
+ 0.2f, 0.1f, 0.4f, 1.6f, 2.4f, 3.3f, 2.1f, 4.2f, 8.4f, 0.0f, 2.1f, 3.2f};
+
+ std::vector bias_data = {
+ -0.5f, 0.6f, 1.2f, 2.1f, 0.5f, 0.7f, 0.2f, 1.2f, 0.5f, 0.4f, 0.3f, 1.2f};
+
+ // Test mask start position > 0.
+ std::vector mask_index_data = {0, 1, 1, 1};
+
+ std::vector output_data = {
+ 3.0146f, 0.1142f, 3.9834f, 5.3394f,
+ 8.69f, -0.13f, 4.25f, 5.65f,
+ 8.69f, -0.13f, 4.25f, 5.65f,
+ -1.4697f, 0.3071f, 4.25f, 5.65f};
+
+ bool use_float16 = true;
+ bool is_unidirectional = true;
+ bool use_past_state = false;
+ int past_sequence_length = 0;
+ const std::vector* past_data = nullptr;
+ const std::vector* present_data = nullptr;
+ RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data,
+ batch_size, sequence_length, hidden_size, number_of_heads,
+ use_float16, is_unidirectional, use_past_state, past_sequence_length, past_data, present_data,
+ AttentionMaskType::MASK_2D_KEY_PADDING, 0 /*input_hidden_size*/, 0 /*max_sequence_length*/,
+ true /*disable_cpu*/, false /*disable_cuda*/, true /*disable_rocm*/, {} /*qkv_sizes*/,
+ {} /*relative_position_bias_data*/, 0 /*kv_sequence_length*/, false /*past_present_share_buffer*/,
+ true /*use_scale*/, true /*use_neox_rotary_embedding*/);
+}
+
TEST(AttentionTest, AttentionMask1DEndNoWord) {
int batch_size = 2;
int sequence_length = 2;
diff --git a/onnxruntime/test/python/transformers/test_parity_neox_attention.py b/onnxruntime/test/python/transformers/test_parity_neox_attention.py
new file mode 100644
index 0000000000..63ed49b901
--- /dev/null
+++ b/onnxruntime/test/python/transformers/test_parity_neox_attention.py
@@ -0,0 +1,310 @@
+# --------------------------------------------------------------------------
+# Copyright 2020 The HuggingFace Inc. team
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at http://www.apache.org/licenses/LICENSE-2.0
+# --------------------------------------------------------------------------
+# Copyright (c) Microsoft Corporation. All rights reserved.
+# Licensed under the MIT License. See License.txt in the project root for
+# license information.
+# -------------------------------------------------------------------------
+
+import numpy as np
+import torch
+from torch import nn
+
+torch.set_printoptions(threshold=10000)
+
+
+def create_neox_attention_graph(
+ batch_size,
+ seq_len,
+ hidden_size,
+ qkv_weight,
+ qkv_bias,
+ num_heads,
+ use_rotary,
+):
+ from onnx import TensorProto, helper
+
+ nodes = [
+ helper.make_node(
+ "Attention",
+ [
+ "input",
+ "weight",
+ "bias",
+ ],
+ ["output"],
+ "NeoXAttention_0",
+ num_heads=num_heads,
+ unidirectional=1,
+ do_rotary=(use_rotary is True),
+ domain="com.microsoft",
+ ),
+ ]
+
+ initializers = [
+ helper.make_tensor("weight", TensorProto.FLOAT, [hidden_size, 3 * hidden_size], qkv_weight.flatten().tolist()),
+ helper.make_tensor("bias", TensorProto.FLOAT, [3 * hidden_size], qkv_bias.flatten().tolist()),
+ ]
+
+ graph = helper.make_graph(
+ nodes,
+ "NeoXAttention_Graph",
+ [
+ helper.make_tensor_value_info("input", TensorProto.FLOAT, [batch_size, seq_len, hidden_size]),
+ ],
+ [
+ helper.make_tensor_value_info("output", TensorProto.FLOAT, [batch_size, seq_len, hidden_size]),
+ ],
+ initializers,
+ )
+
+ model = helper.make_model(graph)
+ return model.SerializeToString()
+
+
+class RotaryEmbedding(torch.nn.Module):
+ def __init__(self, dim, max_position_embeddings, base=10000, device=None):
+ super().__init__()
+ inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float().to(device) / dim))
+ self.register_buffer("inv_freq", inv_freq)
+
+ # Build here to make `torch.jit.trace` work.
+ self.max_seq_len_cached = max_position_embeddings
+ t = torch.arange(self.max_seq_len_cached, device=self.inv_freq.device, dtype=self.inv_freq.dtype)
+ freqs = torch.einsum("i,j->ij", t, self.inv_freq)
+ # Different from paper, but it uses a different permutation in order to obtain the same calculation
+ emb = torch.cat((freqs, freqs), dim=-1)
+ self.cos_cached = emb.cos()[None, None, :, :]
+ self.sin_cached = emb.sin()[None, None, :, :]
+
+ def forward(self, x, seq_len=None):
+ # x: [bs, num_attention_heads, seq_len, head_size]
+ # This `if` block is unlikely to be run after we build sin/cos in `__init__`. Keep the logic here just in case.
+ if seq_len > self.max_seq_len_cached:
+ self.max_seq_len_cached = seq_len
+ t = torch.arange(self.max_seq_len_cached, device=x.device, dtype=self.inv_freq.dtype)
+ freqs = torch.einsum("i,j->ij", t, self.inv_freq)
+ # Different from paper, but it uses a different permutation in order to obtain the same calculation
+ emb = torch.cat((freqs, freqs), dim=-1).to(x.device)
+ self.cos_cached = emb.cos()[None, None, :, :]
+ self.sin_cached = emb.sin()[None, None, :, :]
+ return self.cos_cached[:seq_len, ...].to(x.device), self.sin_cached[:seq_len, ...].to(x.device)
+
+
+def rotate_half(x):
+ """Rotates half the hidden dims of the input."""
+ x1 = x[..., : x.shape[-1] // 2]
+ x2 = x[..., x.shape[-1] // 2 :]
+ return torch.cat((-x2, x1), dim=-1)
+
+
+def apply_rotary_pos_emb(q, k, cos, sin, offset: int = 0):
+ cos = cos[..., offset : q.shape[-2] + offset, :]
+ sin = sin[..., offset : q.shape[-2] + offset, :]
+ q_embed = (q * cos) + (rotate_half(q) * sin)
+ k_embed = (k * cos) + (rotate_half(k) * sin)
+ return q_embed, k_embed
+
+
+class GPTNeoXAttention(nn.Module):
+ def __init__(self, batch_size, seq_len, num_head, hidden_size, use_rotary):
+ super().__init__()
+ self.use_rotary = use_rotary
+ self.num_attention_heads = num_head
+ self.hidden_size = hidden_size
+ self.head_size = self.hidden_size // self.num_attention_heads
+ self.rotary_ndims = int(self.head_size)
+ max_positions = 2048
+ self.register_buffer(
+ "bias",
+ torch.tril(torch.ones((max_positions, max_positions), dtype=torch.uint8)).view(
+ 1, 1, max_positions, max_positions
+ ),
+ )
+ self.register_buffer("masked_bias", torch.tensor(-1e9))
+ self.rotary_emb = RotaryEmbedding(self.rotary_ndims, 2048, 10000)
+ self.norm_factor = torch.sqrt(torch.tensor(self.head_size, dtype=torch.float32)).to(torch.get_default_dtype())
+ self.query_key_value = nn.Linear(hidden_size, 3 * hidden_size)
+
+ # self.query_key_value.weight.data.copy_(torch.tensor(np.ones((3 * hidden_size, hidden_size))))
+ # self.query_key_value.bias.data.copy_(torch.tensor(np.zeros((3 * hidden_size))))
+
+ self.onnx_graph = create_neox_attention_graph(
+ batch_size,
+ seq_len,
+ self.hidden_size,
+ self.query_key_value.weight.reshape(self.num_attention_heads, 3, -1)
+ .transpose(0, 1)
+ .reshape(3 * self.hidden_size, -1)
+ .transpose(0, 1),
+ self.query_key_value.bias.reshape(self.num_attention_heads, 3, -1).transpose(0, 1).reshape(-1),
+ self.num_attention_heads,
+ self.use_rotary,
+ )
+
+ @classmethod
+ def _merge_heads(cls, tensor, num_attention_heads, attn_head_size):
+ """
+ Merges attn_head_size dim and num_attn_heads dim into hidden dim
+ """
+ # tensor [bs, num_attention_heads, seq_len, attn_head_size]
+ tensor = tensor.permute(0, 2, 1, 3).contiguous()
+ # -> [bs, seq_len, num_attention_heads, attn_head_size]
+ tensor = tensor.view(tensor.size(0), tensor.size(1), num_attention_heads * attn_head_size)
+ # -> [bs, seq_len, hidden_size]
+ return tensor
+
+ def _attn(self, query, key, value, attention_mask=None, head_mask=None):
+ # q, k, v: [bs, num_attention_heads, seq_len, attn_head_size]
+ # compute causal mask from causal mask buffer
+ batch_size, num_attention_heads, query_length, attn_head_size = query.size()
+ key_length = key.size(-2)
+
+ causal_mask = self.bias[:, :, key_length - query_length : key_length, :key_length].bool()
+
+ query = query.view(batch_size * num_attention_heads, query_length, attn_head_size)
+ key = key.view(batch_size * num_attention_heads, key_length, attn_head_size)
+ attn_scores = torch.zeros(
+ batch_size * num_attention_heads,
+ query_length,
+ key_length,
+ dtype=query.dtype,
+ device=key.device,
+ )
+ attn_scores = torch.baddbmm(
+ attn_scores,
+ query,
+ key.transpose(1, 2),
+ beta=1.0,
+ alpha=(torch.tensor(1.0, dtype=self.norm_factor.dtype, device=self.norm_factor.device) / self.norm_factor),
+ )
+ attn_scores = attn_scores.view(batch_size, num_attention_heads, query_length, key_length)
+
+ mask_value = torch.finfo(attn_scores.dtype).min
+ # Need to be a tensor, otherwise we get error: `RuntimeError: expected scalar type float but found double`.
+ # Need to be on the same device, otherwise `RuntimeError: ..., x and y to be on the same device`
+ mask_value = torch.tensor(mask_value, dtype=attn_scores.dtype).to(attn_scores.device)
+ attn_scores = torch.where(causal_mask, attn_scores, mask_value)
+
+ if attention_mask is not None:
+ # Apply the attention mask
+ attn_scores = attn_scores + attention_mask
+
+ attn_weights = nn.functional.softmax(attn_scores, dim=-1)
+ attn_weights = attn_weights.to(value.dtype)
+
+ # Mask heads if we want to
+ if head_mask is not None:
+ attn_weights = attn_weights * head_mask
+
+ attn_output = torch.matmul(attn_weights, value)
+ return attn_output, attn_weights
+
+ def onnx_forward(
+ self,
+ hidden_states,
+ ):
+ ort_inputs = {
+ "input": np.ascontiguousarray(hidden_states.cpu().numpy()),
+ }
+
+ from onnxruntime import InferenceSession, SessionOptions
+
+ sess_options = SessionOptions()
+ ort_session = InferenceSession(self.onnx_graph, sess_options, providers=["CUDAExecutionProvider"])
+ ort_output = ort_session.run(None, ort_inputs)
+
+ output = torch.tensor(ort_output)
+
+ return output
+
+ def torch_forward(
+ self,
+ hidden_states,
+ attention_mask=None,
+ head_mask=None,
+ layer_past=None,
+ use_cache=False,
+ output_attentions=False,
+ ):
+ has_layer_past = layer_past is not None
+
+ # Compute QKV
+ # Attention heads [batch, seq_len, hidden_size]
+ # --> [batch, seq_len, (np * 3 * head_size)]
+ qkv = self.query_key_value(hidden_states)
+
+ # [batch, seq_len, (num_heads * 3 * head_size)]
+ # --> [batch, seq_len, num_heads, 3 * head_size]
+ new_qkv_shape = qkv.size()[:-1] + (self.num_attention_heads, 3 * self.head_size)
+ qkv = qkv.view(*new_qkv_shape)
+
+ # [batch, seq_len, num_attention_heads, 3 * head_size] --> 3 [batch, num_attention_heads, seq_len, head_size]
+ query = qkv[..., : self.head_size].permute(0, 2, 1, 3)
+ key = qkv[..., self.head_size : 2 * self.head_size].permute(0, 2, 1, 3)
+ value = qkv[..., 2 * self.head_size :].permute(0, 2, 1, 3)
+
+ if self.use_rotary:
+ # Compute rotary embeddings on rotary_ndims
+ query_rot = query[..., : self.rotary_ndims]
+ query_pass = query[..., self.rotary_ndims :]
+ key_rot = key[..., : self.rotary_ndims]
+ key_pass = key[..., self.rotary_ndims :]
+
+ # Compute token offset for rotary embeddings (when decoding)
+ seq_len = key.shape[-2]
+ offset = 0
+ if has_layer_past:
+ offset = layer_past[0].shape[-2]
+ seq_len += offset
+ cos, sin = self.rotary_emb(value, seq_len=seq_len)
+ query, key = apply_rotary_pos_emb(query_rot, key_rot, cos, sin, offset=offset)
+ query = torch.cat((query, query_pass), dim=-1)
+ key = torch.cat((key, key_pass), dim=-1)
+ # print("query", query.shape, query)
+ # print("key", key.shape, key)
+ # print("value", value.shape, value)
+
+ # Cache QKV values
+ if has_layer_past:
+ past_key = layer_past[0]
+ past_value = layer_past[1]
+ key = torch.cat((past_key, key), dim=-2)
+ value = torch.cat((past_value, value), dim=-2)
+ present = (key, value) if use_cache else None
+
+ # Compute attention
+ attn_output, _ = self._attn(query, key, value, attention_mask, head_mask)
+
+ # Reshape outputs
+ attn_output = self._merge_heads(attn_output, self.num_attention_heads, self.head_size)
+
+ return attn_output
+
+
+if __name__ == "__main__":
+ for batch_size in [1, 2, 4, 8]:
+ for seq_len in [32, 128, 512, 1024, 2048]:
+ for num_head in [12]:
+ for hidden_size in [768]:
+ attn = GPTNeoXAttention(batch_size, seq_len, num_head, hidden_size, use_rotary=True)
+
+ hidden_states = torch.normal(mean=0.5, std=0.1, size=(batch_size, seq_len, hidden_size)).to(
+ torch.float32
+ )
+
+ torch_output = attn.torch_forward(hidden_states)
+ ort_output = attn.onnx_forward(hidden_states)
+ print(
+ "Parity check with shape BNSH = ({},{},{},{})".format(
+ batch_size, seq_len, num_head, hidden_size
+ )
+ )
+ if torch.allclose(torch_output, ort_output, atol=1e-6):
+ print("Success!")
+ else:
+ print("Failure!")