diff --git a/docs/ContribOperators.md b/docs/ContribOperators.md
old mode 100755
new mode 100644
index 99ad6a6d0f..824df77282
--- a/docs/ContribOperators.md
+++ b/docs/ContribOperators.md
@@ -4417,9 +4417,9 @@ This version of the operator has been available since version 1 of the 'com.micr
- input : T
-- 3D input tensor with shape (batch_size, sequence_length, hidden_size)
+- 3D input tensor with shape (batch_size, sequence_length, hidden_size)Or 2D input tensor with shape (token_count, hidden_size)
- skip : T
-- 3D skip tensor with shape (batch_size, sequence_length, hidden_size)
+- 3D input tensor with shape (batch_size, sequence_length, hidden_size)Or 2D input tensor with shape (token_count, hidden_size)
- gamma : T
- 1D input tensor with shape (hidden_size)
- bias (optional) : T
@@ -4430,13 +4430,13 @@ This version of the operator has been available since version 1 of the 'com.micr
- output : T
-- 3D output tensor with shape (batch_size, sequence_length, hidden_size)
+- 3D output tensor with shape (batch_size, sequence_length, hidden_size)Or 2D output tensor with shape (token_count, hidden_size)
- mean (optional) : U
- Saved mean used during training to speed up gradient computation
- inv_std_var (optional) : U
- Saved inverse standard variance used during training to speed up gradient computation.
- input_skip_bias_sum (optional) : T
-- Sum of the input and skip inputs (and bias if it exists) with shape (batch_size, sequence_length, hidden_size).
+- Sum of the input and skip inputs (and bias if it exists)with shape (batch_size, sequence_length, hidden_size) or (token_count, hidden_size).
#### Type Constraints
diff --git a/onnxruntime/contrib_ops/cpu/skip_layer_norm.cc b/onnxruntime/contrib_ops/cpu/skip_layer_norm.cc
index a394ebb4ce..486f73f261 100644
--- a/onnxruntime/contrib_ops/cpu/skip_layer_norm.cc
+++ b/onnxruntime/contrib_ops/cpu/skip_layer_norm.cc
@@ -44,11 +44,14 @@ Status SkipLayerNorm::Compute(OpKernelContext* p_ctx) const {
Tensor* skip_input_bias_add_output = p_ctx->Output(3, input->Shape());
const auto& input_dims = input->Shape().GetDims();
- if (input_dims.size() != 3) {
+ size_t input_dims_size = input_dims.size();
+ if (input_dims_size != 3 && input_dims_size != 2) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
- "input is expected to have 3 dimensions, got ", input_dims.size());
+ "input is expected to have 3 or 2 dimensions, got ", input_dims_size);
}
+ int hidden_size = static_cast(input_dims[input_dims_size - 1]);
+
if (input->Shape() != skip->Shape()) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"skip is expected to have same shape as input");
@@ -59,7 +62,7 @@ Status SkipLayerNorm::Compute(OpKernelContext* p_ctx) const {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"gamma is expected to have 1 dimension, got ", gamma_dims.size());
}
- if (gamma_dims[0] != input_dims[2]) {
+ if (gamma_dims[0] != hidden_size) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"Last dimension of gamma and input does not match");
}
@@ -70,7 +73,7 @@ Status SkipLayerNorm::Compute(OpKernelContext* p_ctx) const {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"beta is expected to have 1 dimension, got ", beta_dims.size());
}
- if (beta_dims[0] != input_dims[2]) {
+ if (beta_dims[0] != hidden_size) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"Last dimension of beta and input does not match");
}
@@ -82,16 +85,14 @@ Status SkipLayerNorm::Compute(OpKernelContext* p_ctx) const {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"bias is expected to have 1 dimension, got ", bias_dims.size());
}
- if (bias_dims[0] != input_dims[2]) {
+ if (bias_dims[0] != hidden_size) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"Last dimension of bias and input does not match");
}
}
- int64_t batch_size = input_dims[0];
- int64_t sequence_length = input_dims[1];
- int64_t hidden_size = input_dims[2];
- int64_t task_count = batch_size * sequence_length;
+
+ int64_t task_count = input->Shape().SizeToDimension(input_dims_size - 1);
const T* input_data = input->Data();
const T* skip_data = skip->Data();
diff --git a/onnxruntime/contrib_ops/cuda/bert/skip_layer_norm.cc b/onnxruntime/contrib_ops/cuda/bert/skip_layer_norm.cc
index d2f7d974be..4cf45b5d33 100644
--- a/onnxruntime/contrib_ops/cuda/bert/skip_layer_norm.cc
+++ b/onnxruntime/contrib_ops/cuda/bert/skip_layer_norm.cc
@@ -65,17 +65,20 @@ Status SkipLayerNorm::ComputeInternal(OpKernelContext* ctx) const
}
const auto& input_dims = input->Shape().GetDims();
- if (input_dims.size() != 3) {
+ size_t input_dims_size = input_dims.size();
+ if (input_dims_size != 3 && input_dims_size != 2) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
- "input is expected to have 3 dimensions, got ", input_dims.size());
+ "input is expected to have 3 or 2 dimensions, got ", input_dims_size);
}
+ int hidden_size = static_cast(input_dims[input_dims_size - 1]);
+
const auto& gamma_dims = gamma->Shape().GetDims();
if (gamma_dims.size() != 1) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"gamma is expected to have 1 dimension, got ", gamma_dims.size());
}
- if (gamma_dims[0] != input_dims[2]) {
+ if (gamma_dims[0] != hidden_size) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"Last dimension of gamma and input does not match");
}
@@ -87,7 +90,7 @@ Status SkipLayerNorm::ComputeInternal(OpKernelContext* ctx) const
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"beta is expected to have 1 dimension, got ", beta_dims.size());
}
- if (beta_dims[0] != input_dims[2]) {
+ if (beta_dims[0] != hidden_size) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"Last dimension of beta and input does not match");
}
@@ -100,16 +103,13 @@ Status SkipLayerNorm::ComputeInternal(OpKernelContext* ctx) const
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"bias is expected to have 1 dimension, got ", bias_dims.size());
}
- if (bias_dims[0] != input_dims[2]) {
+ if (bias_dims[0] != hidden_size) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"Last dimension of bias and input does not match");
}
}
- int sequence_length = gsl::narrow_cast(input_dims[1]);
- int hidden_size = gsl::narrow_cast(input_dims[2]);
- int row_count = gsl::narrow_cast(input_dims[0] * sequence_length);
-
+ int row_count = gsl::narrow(input->Shape().SizeToDimension(input_dims_size - 1));
typedef typename ToCudaType::MappedType CudaT;
HostApplyLayerNorm(
GetDeviceProp(),
diff --git a/onnxruntime/contrib_ops/rocm/bert/skip_layer_norm.cc b/onnxruntime/contrib_ops/rocm/bert/skip_layer_norm.cc
index 24dbb87b50..5b0f57c8f7 100644
--- a/onnxruntime/contrib_ops/rocm/bert/skip_layer_norm.cc
+++ b/onnxruntime/contrib_ops/rocm/bert/skip_layer_norm.cc
@@ -57,17 +57,20 @@ Status SkipLayerNorm::ComputeInternal(OpKernelContext* ctx) const {
}
const auto& input_dims = input->Shape().GetDims();
- if (input_dims.size() != 3) {
+ size_t input_dims_size = input_dims.size();
+ if (input_dims_size != 3 && input_dims_size != 2) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
- "input is expected to have 3 dimensions, got ", input_dims.size());
+ "input is expected to have 3 or 2 dimensions, got ", input_dims_size);
}
+ int hidden_size = static_cast(input_dims[input_dims_size - 1]);
+
const auto& gamma_dims = gamma->Shape().GetDims();
if (gamma_dims.size() != 1) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"gamma is expected to have 1 dimension, got ", gamma_dims.size());
}
- if (gamma_dims[0] != input_dims[2]) {
+ if (gamma_dims[0] != hidden_size) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"Last dimension of gamma and input does not match");
}
@@ -78,7 +81,7 @@ Status SkipLayerNorm::ComputeInternal(OpKernelContext* ctx) const {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"beta is expected to have 1 dimension, got ", beta_dims.size());
}
- if (beta_dims[0] != input_dims[2]) {
+ if (beta_dims[0] != hidden_size) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"Last dimension of beta and input does not match");
}
@@ -90,15 +93,13 @@ Status SkipLayerNorm::ComputeInternal(OpKernelContext* ctx) const {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"bias is expected to have 1 dimension, got ", bias_dims.size());
}
- if (bias_dims[0] != input_dims[2]) {
+ if (bias_dims[0] != hidden_size) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"Last dimension of bias and input does not match");
}
}
- int sequence_length = static_cast(input_dims[1]);
- int hidden_size = static_cast(input_dims[2]);
- int64_t element_count = input_dims[0] * sequence_length * hidden_size;
+ int64_t element_count = input->Shape().Size();
typedef typename ToHipType::MappedType HipT;
return LaunchSkipLayerNormKernel(
diff --git a/onnxruntime/core/graph/contrib_ops/bert_defs.cc b/onnxruntime/core/graph/contrib_ops/bert_defs.cc
index e4e0f53886..174dde6358 100644
--- a/onnxruntime/core/graph/contrib_ops/bert_defs.cc
+++ b/onnxruntime/core/graph/contrib_ops/bert_defs.cc
@@ -814,14 +814,46 @@ ONNX_MS_OPERATOR_SET_SCHEMA(
OpSchema()
.SetDoc("Skip and Root Mean Square Layer Normalization")
.Attr("epsilon", "The epsilon value to use to avoid division by zero.", AttributeProto::FLOAT, kDefaultSkipLayerNormEpsilon)
- .Input(0, "input", "3D input tensor with shape (batch_size, sequence_length, hidden_size)", "T")
- .Input(1, "skip", "3D skip tensor with shape (batch_size, sequence_length, hidden_size)", "T")
- .Input(2, "gamma", "1D input tensor with shape (hidden_size)", "T")
- .Input(3, "bias", "1D bias tensor with shape (hidden_size", "T", OpSchema::Optional)
- .Output(0, "output", "3D output tensor with shape (batch_size, sequence_length, hidden_size)", "T")
- .Output(1, "mean", "Saved mean used during training to speed up gradient computation", "U", OpSchema::Optional)
- .Output(2, "inv_std_var", "Saved inverse standard variance used during training to speed up gradient computation.", "U", OpSchema::Optional)
- .Output(3, "input_skip_bias_sum", "Sum of the input and skip inputs (and bias if it exists) with shape (batch_size, sequence_length, hidden_size).", "T", OpSchema::Optional)
+ .Input(0,
+ "input",
+ "3D input tensor with shape (batch_size, sequence_length, hidden_size)"
+ "Or 2D input tensor with shape (token_count, hidden_size)",
+ "T")
+ .Input(1,
+ "skip",
+ "3D input tensor with shape (batch_size, sequence_length, hidden_size)"
+ "Or 2D input tensor with shape (token_count, hidden_size)",
+ "T")
+ .Input(2,
+ "gamma",
+ "1D input tensor with shape (hidden_size)",
+ "T")
+ .Input(3,
+ "bias",
+ "1D bias tensor with shape (hidden_size",
+ "T",
+ OpSchema::Optional)
+ .Output(0,
+ "output",
+ "3D output tensor with shape (batch_size, sequence_length, hidden_size)"
+ "Or 2D output tensor with shape (token_count, hidden_size)",
+ "T")
+ .Output(1,
+ "mean",
+ "Saved mean used during training to speed up gradient computation",
+ "U",
+ OpSchema::Optional)
+ .Output(2,
+ "inv_std_var",
+ "Saved inverse standard variance used during training to speed up gradient computation.",
+ "U",
+ OpSchema::Optional)
+ .Output(3,
+ "input_skip_bias_sum",
+ "Sum of the input and skip inputs (and bias if it exists)"
+ "with shape (batch_size, sequence_length, hidden_size) or (token_count, hidden_size).",
+ "T",
+ OpSchema::Optional)
.TypeConstraint("T", {"tensor(float)", "tensor(float16)"}, "Constrain input and output types to float or half tensors.")
.TypeConstraint("U", {"tensor(float)"}, "Constrain mean and inv_std_var to float tensors.")
.TypeAndShapeInferenceFunction(ONNX_NAMESPACE::propagateShapeAndTypeFromFirstInput));
diff --git a/onnxruntime/core/providers/dnnl/subgraph/dnnl_layernorm.cc b/onnxruntime/core/providers/dnnl/subgraph/dnnl_layernorm.cc
index 1e21a95598..d9c8e8a6ef 100644
--- a/onnxruntime/core/providers/dnnl/subgraph/dnnl_layernorm.cc
+++ b/onnxruntime/core/providers/dnnl/subgraph/dnnl_layernorm.cc
@@ -24,13 +24,13 @@ Layer Norm:
+-----------+
(X) ---------->+ +----------> (Y)
- | |
+ | |
(G) ---------->+ LayerNorm +----------> (M)
- | |
-(B) ---------->+ +----------> (I)
- +-----------+
-
-
+ | |
+(B) ---------->+ +----------> (I)
+ +-----------+
+
+
Skip Layer Norm:
Inputs:
@@ -56,7 +56,7 @@ Skip Layer Norm:
(E) ------->+ +----------> (I)
| |
+-----------+
-
+
Attributes (epsilon)
*/
void DnnlLayerNorm::CreatePrimitive(DnnlSubgraphPrimitive& sp, DnnlNode& node) {
@@ -82,7 +82,7 @@ void DnnlLayerNorm::CreatePrimitive(DnnlSubgraphPrimitive& sp, DnnlNode& node) {
// This contains the layer norm op and its parameters
ln_components op_comps;
if (node.OpType() == "SkipLayerNormalization") {
-
+
// Check if shift is available
shift_exists = node.Input(IN_BETA).Exists();
@@ -258,11 +258,6 @@ void DnnlLayerNorm::ValidateDims(DnnlSubgraphPrimitive& sp, DnnlNode& node) {
// define gamma and shift input position, depending on the operation
int gamma_pos, shift_pos;
if (node.OpType() == "SkipLayerNormalization") {
- // For SkipLayerNorm the spec defines the input as a 3D tensor
- if (input_dims_size != 3) {
- // We support 2D arrays but the expected is 3D
- ORT_THROW("Input tensor is expected to have 3 dimensions, got ", input_dims_size);
- }
// Get skip and evaluate
auto skip_dims = sp.GetMemory(node.Input(IN_SKIP)).get_desc().get_dims();
@@ -343,4 +338,4 @@ dnnl::memory DnnlLayerNorm::CastAndTransformMemory(DnnlSubgraphPrimitive& sp, dn
}
} // namespace ort_dnnl
-} // namespace onnxruntime
\ No newline at end of file
+} // namespace onnxruntime
diff --git a/onnxruntime/python/tools/symbolic_shape_infer.py b/onnxruntime/python/tools/symbolic_shape_infer.py
index fe8b1951e0..bb7f3ecaab 100755
--- a/onnxruntime/python/tools/symbolic_shape_infer.py
+++ b/onnxruntime/python/tools/symbolic_shape_infer.py
@@ -190,6 +190,9 @@ class SymbolicShapeInference:
"Neg": self._infer_symbolic_compute_ops,
# contrib ops:
"Attention": self._infer_Attention,
+ "PackedAttention": self._infer_PackedAttention,
+ "RemovePadding": self._infer_RemovePadding,
+ "RestorePadding": self._infer_RestorePadding,
"BiasGelu": self._infer_BiasGelu,
"MultiHeadAttention": self._infer_MultiHeadAttention,
"EmbedLayerNormalization": self._infer_EmbedLayerNormalization,
@@ -445,9 +448,12 @@ class SymbolicShapeInference:
"LayerNormalization",
"LongformerAttention",
"RelativePositionBias",
+ "RemovePadding",
+ "RestorePadding",
"SimplifiedLayerNormalization",
"SkipLayerNormalization",
"SkipSimplifiedLayerNormalization",
+ "PackedAttention",
"PythonOp",
"MultiHeadAttention",
"GroupNorm",
@@ -2097,6 +2103,54 @@ class SymbolicShapeInference:
vi = self.known_vi_[node.output[1]]
vi.CopyFrom(helper.make_tensor_value_info(vi.name, output_dtype, past_shape))
+ def _infer_PackedAttention(self, node): # noqa: N802
+ shape = self._get_shape(node, 0)
+ shape_weights = self._get_shape(node, 1)
+ shape_bias = self._try_get_shape(node, 2)
+ if shape_bias is not None:
+ assert len(shape_bias) == 1
+ tripled_hidden_size = shape_bias[0] if shape_bias is not None else shape_weights[1]
+ if shape and len(shape) == 2:
+ qkv_hidden_sizes_attr = get_attribute(node, "qkv_hidden_sizes")
+ if qkv_hidden_sizes_attr is not None:
+ assert len(qkv_hidden_sizes_attr) == 3
+ shape[1] = int(qkv_hidden_sizes_attr[2])
+ elif isinstance(tripled_hidden_size, int):
+ shape[1] = int(tripled_hidden_size / 3)
+ output_dtype = self.known_vi_[node.input[0]].type.tensor_type.elem_type
+ vi = self.known_vi_[node.output[0]]
+ vi.CopyFrom(helper.make_tensor_value_info(node.output[0], output_dtype, shape))
+
+ def _infer_RemovePadding(self, node): # noqa: N802
+ shape = self._get_shape(node, 0)
+ if shape and len(shape) == 3:
+ output_dtype = self.known_vi_[node.input[0]].type.tensor_type.elem_type
+ vi = self.known_vi_[node.output[0]]
+ vi.CopyFrom(helper.make_tensor_value_info(node.output[0], output_dtype, ["token_count", shape[2]]))
+
+ vi_token_offset = self.known_vi_[node.output[1]]
+ vi_token_offset.CopyFrom(
+ helper.make_tensor_value_info(node.output[1], onnx.TensorProto.INT32, [shape[0], shape[1]])
+ )
+
+ vi_cumulated_seq_len = self.known_vi_[node.output[2]]
+ vi_cumulated_seq_len.CopyFrom(
+ helper.make_tensor_value_info(node.output[2], onnx.TensorProto.INT32, ["batch_size + 1"])
+ )
+
+ vi_max_seq_len = self.known_vi_[node.output[3]]
+ vi_max_seq_len.CopyFrom(helper.make_tensor_value_info(node.output[3], onnx.TensorProto.INT32, [1]))
+
+ def _infer_RestorePadding(self, node): # noqa: N802
+ shape_input = self._get_shape(node, 0)
+ shape_token_offset = self._get_shape(node, 1)
+ if shape_input and len(shape_input) == 2 and shape_token_offset and len(shape_token_offset) == 2:
+ output_dtype = self.known_vi_[node.input[0]].type.tensor_type.elem_type
+ vi = self.known_vi_[node.output[0]]
+
+ output_shape = [shape_token_offset[0], shape_token_offset[1], shape_input[1]]
+ vi.CopyFrom(helper.make_tensor_value_info(node.output[0], output_dtype, output_shape))
+
def _infer_BiasGelu(self, node): # noqa: N802
self._propagate_shape_and_type(node)
diff --git a/onnxruntime/python/tools/transformers/constants.py b/onnxruntime/python/tools/transformers/constants.py
new file mode 100644
index 0000000000..9f12d4de5a
--- /dev/null
+++ b/onnxruntime/python/tools/transformers/constants.py
@@ -0,0 +1,28 @@
+# -------------------------------------------------------------------------
+# Copyright (c) Microsoft Corporation. All rights reserved.
+# Licensed under the MIT License.
+# --------------------------------------------------------------------------
+
+
+class Operators:
+ ATTENTION = "Attention"
+ LAYERNORM = "LayerNormalization"
+ PACKEDATTENTION = "PackedAttention"
+ REMOVEPADDING = "RemovePadding"
+ RESTOREPADDING = "RestorePadding"
+ SKIPLAYERNORM = "SkipLayerNormalization"
+
+
+class AttentionInputIDs:
+ INPUT = 0
+ WEIGHTS = 1
+ BIAS = 2
+ MASK_INDEX = 3
+ PAST = 4
+ RELATIVE_POSITION_BIAS = 5
+ PAST_SEQUENCE_LENGTH = 6
+
+
+class AttentionOutputIDs:
+ OUTPUT = 0
+ PRESENT = 1
diff --git a/onnxruntime/python/tools/transformers/convert_to_packing_mode.py b/onnxruntime/python/tools/transformers/convert_to_packing_mode.py
new file mode 100644
index 0000000000..16d812169b
--- /dev/null
+++ b/onnxruntime/python/tools/transformers/convert_to_packing_mode.py
@@ -0,0 +1,241 @@
+# -------------------------------------------------------------------------
+# Copyright (c) Microsoft Corporation. All rights reserved.
+# Licensed under the MIT License.
+# --------------------------------------------------------------------------
+
+import argparse
+import logging
+import os
+from typing import List, Union
+
+import coloredlogs
+from constants import AttentionInputIDs, AttentionOutputIDs, Operators
+from onnx import helper, load_model
+from onnx_model import NodeProto, OnnxModel
+from shape_infer_helper import SymbolicShapeInferenceHelper
+
+logger = logging.getLogger(__name__)
+
+
+class PackingMode:
+ def __init__(
+ self,
+ model: OnnxModel,
+ ):
+ self.model: OnnxModel = model
+ self.nodes_to_remove: List = []
+ self.nodes_to_add: List = []
+ self.prune_graph: bool = False
+ self.node_name_to_graph_name: dict = {}
+ self.this_graph_name: str = self.model.model.graph.name
+ self.attention_nodes = self.model.get_nodes_by_op_type(Operators.ATTENTION)
+
+ def _try_getting_attention_mask(self) -> Union[str, None]:
+ first_attention_node = self._try_getting_first_attention()
+ # check if attention has mask
+ if not first_attention_node or len(first_attention_node.input) <= AttentionInputIDs.MASK_INDEX:
+ return None
+
+ attention_mask = first_attention_node.input[AttentionInputIDs.MASK_INDEX]
+
+ # check if all attention nodes have same mask
+ for node in self.attention_nodes:
+ if (
+ len(node.input) <= AttentionInputIDs.MASK_INDEX
+ or node.input[AttentionInputIDs.MASK_INDEX] != attention_mask
+ ):
+ return None
+
+ return attention_mask
+
+ def _try_getting_first_attention(self) -> Union[NodeProto, None]:
+ if len(self.attention_nodes) <= 0:
+ return None
+
+ return self.attention_nodes[0]
+
+ def _try_getting_last_layernorm(self) -> Union[NodeProto, None]:
+ last_layernorm_node = None
+ for node in self.model.nodes():
+ if node.op_type == Operators.LAYERNORM or node.op_type == Operators.SKIPLAYERNORM:
+ last_layernorm_node = node
+ return last_layernorm_node
+
+ def _are_attentions_supportted(self) -> bool:
+ for node in self.attention_nodes:
+ if OnnxModel.get_node_attribute(node, "past_present_share_buffer") is not None:
+ return False
+ if OnnxModel.get_node_attribute(node, "do_rotary") is not None:
+ return False
+ unidirection_attr = OnnxModel.get_node_attribute(node, "unidirectional")
+ if unidirection_attr is not None and unidirection_attr != 0:
+ return False
+ if len(node.input) > AttentionInputIDs.PAST and not node.input[AttentionInputIDs.PAST]:
+ return False
+ if (
+ len(node.input) > AttentionInputIDs.PAST_SEQUENCE_LENGTH
+ and not node.input[AttentionInputIDs.PAST_SEQUENCE_LENGTH]
+ ):
+ return False
+ return True
+
+ def _insert_removepadding_node(self, inputs: List[str], outputs: List[str]) -> None:
+ new_node = helper.make_node(
+ Operators.REMOVEPADDING,
+ inputs=inputs,
+ outputs=outputs,
+ name=self.model.create_node_name(Operators.REMOVEPADDING),
+ )
+
+ new_node.domain = "com.microsoft"
+ self.nodes_to_add.append(new_node)
+ self.node_name_to_graph_name[new_node.name] = self.this_graph_name
+
+ def _insert_restorepadding_node(self, inputs: List[str], outputs: List[str]) -> None:
+ new_node = helper.make_node(
+ Operators.RESTOREPADDING,
+ inputs=inputs,
+ outputs=outputs,
+ name=self.model.create_node_name(Operators.RESTOREPADDING),
+ )
+
+ new_node.domain = "com.microsoft"
+ self.nodes_to_add.append(new_node)
+ self.node_name_to_graph_name[new_node.name] = self.this_graph_name
+
+ def _replace_attention_with_packing_attention(self, token_offset: str, cumulative_sequence_length: str) -> None:
+ for attention in self.attention_nodes:
+ packed_attention = helper.make_node(
+ Operators.PACKEDATTENTION,
+ inputs=[
+ attention.input[AttentionInputIDs.INPUT],
+ attention.input[AttentionInputIDs.WEIGHTS],
+ attention.input[AttentionInputIDs.BIAS],
+ token_offset,
+ cumulative_sequence_length,
+ attention.input[AttentionInputIDs.RELATIVE_POSITION_BIAS]
+ if len(attention.input) > AttentionInputIDs.RELATIVE_POSITION_BIAS
+ else "",
+ ],
+ outputs=[attention.output[AttentionOutputIDs.OUTPUT]],
+ name=self.model.create_node_name(Operators.PACKEDATTENTION),
+ )
+
+ attributes = []
+ for attr in attention.attribute:
+ if attr.name in ["num_heads", "qkv_hidden_sizes", "scale"]:
+ attributes.append(attr)
+
+ packed_attention.attribute.extend(attributes)
+ packed_attention.domain = "com.microsoft"
+ self.nodes_to_add.append(packed_attention)
+ self.nodes_to_remove.append(attention)
+ self.node_name_to_graph_name[packed_attention.name] = self.this_graph_name
+
+ def convert(self, use_symbolic_shape_infer: bool = True) -> None:
+ logger.debug("start converting to packing model...")
+ if not self._are_attentions_supportted():
+ return
+
+ attention_mask = self._try_getting_attention_mask()
+ if not attention_mask:
+ return
+
+ first_attention_node = self._try_getting_first_attention()
+ last_layernorm_node = self._try_getting_last_layernorm()
+ if not last_layernorm_node:
+ return
+
+ # insert RemovePadding
+ first_attention_input = first_attention_node.input[AttentionInputIDs.INPUT]
+ input_to_remove_padding = first_attention_input
+ output_without_padding = first_attention_input + "_no_padding"
+ token_offset = first_attention_input + "_token_offset"
+ cumulated_seq_len = first_attention_input + "_cumulated_seq_len"
+ max_seq_len = first_attention_input + "_max_seq_len"
+ self._insert_removepadding_node(
+ [input_to_remove_padding, attention_mask],
+ [output_without_padding, token_offset, cumulated_seq_len, max_seq_len],
+ )
+ self.model.replace_input_of_all_nodes(input_to_remove_padding, output_without_padding)
+ logger.debug("inserted RemovePadding before Attention")
+
+ # insert RestorePadding
+ restorepadding_input = last_layernorm_node.output[0] + "_restore_input"
+ self._insert_restorepadding_node([restorepadding_input, token_offset], [last_layernorm_node.output[0]])
+ self.model.replace_output_of_all_nodes(last_layernorm_node.output[0], restorepadding_input)
+ logger.debug(f"inserted RestorePadding after last {last_layernorm_node.op_type} layer")
+
+ # insert PackingAttention
+ self._replace_attention_with_packing_attention(token_offset, cumulated_seq_len)
+ logger.debug("replaced Attention with PackedAttention")
+
+ self.model.remove_nodes(self.nodes_to_remove)
+ self.model.add_nodes(self.nodes_to_add, self.node_name_to_graph_name)
+
+ if self.prune_graph:
+ self.model.prune_graph()
+ elif self.nodes_to_remove or self.nodes_to_add:
+ self.model.update_graph()
+ self.model.clean_shape_infer()
+ if use_symbolic_shape_infer:
+ # Use symbolic shape inference since custom operators (like Gelu, SkipLayerNormalization etc)
+ # are not recognized by onnx shape inference.
+ shape_infer_helper = SymbolicShapeInferenceHelper(self.model.model, verbose=0)
+ inferred_model = shape_infer_helper.infer_shapes(self.model.model, auto_merge=True, guess_output_rank=False)
+ if inferred_model:
+ self.model.model = inferred_model
+
+
+def _parse_arguments():
+ parser = argparse.ArgumentParser(
+ description="Convert to packing mode tool for ONNX Runtime." "It converts BERT like model to use packing mode."
+ )
+ parser.add_argument("--input", required=True, type=str, help="input onnx model path")
+
+ parser.add_argument("--output", required=True, type=str, help="optimized onnx model path")
+
+ parser.add_argument("--verbose", required=False, action="store_true", help="show debug information.")
+ parser.set_defaults(verbose=False)
+
+ parser.add_argument(
+ "--use_external_data_format",
+ required=False,
+ action="store_true",
+ help="use external data format to store large model (>2GB)",
+ )
+ parser.set_defaults(use_external_data_format=False)
+
+ args = parser.parse_args()
+
+ return args
+
+
+def _setup_logger(verbose):
+ if verbose:
+ coloredlogs.install(
+ level="DEBUG",
+ fmt="[%(filename)s:%(lineno)s - %(funcName)20s()] %(message)s",
+ )
+ else:
+ coloredlogs.install(fmt="%(funcName)20s: %(message)s")
+
+
+def main():
+ args = _parse_arguments()
+
+ _setup_logger(args.verbose)
+
+ logger.debug("arguments:{args}")
+
+ if os.path.realpath(args.input) == os.path.realpath(args.output):
+ logger.warning("Specified the same input and output path. Note that this may overwrite the original model")
+
+ model = load_model(args.input)
+ packing_mode = PackingMode(OnnxModel(model))
+ packing_mode.convert()
+ packing_mode.model.save_model_to_file(args.output, use_external_data_format=args.use_external_data_format)
+
+
+if __name__ == "__main__":
+ main()
diff --git a/onnxruntime/python/tools/transformers/onnx_model.py b/onnxruntime/python/tools/transformers/onnx_model.py
index aab2358e2b..bf199887e5 100644
--- a/onnxruntime/python/tools/transformers/onnx_model.py
+++ b/onnxruntime/python/tools/transformers/onnx_model.py
@@ -1124,3 +1124,6 @@ class OnnxModel:
for value_info in self.model.graph.value_info:
if value_info.name not in excluded:
value_info.name = prefix + value_info.name
+
+ def clean_shape_infer(self):
+ self.model.graph.ClearField("value_info")
diff --git a/onnxruntime/python/tools/transformers/onnx_model_bert.py b/onnxruntime/python/tools/transformers/onnx_model_bert.py
index c8288b4b15..a3c0470321 100644
--- a/onnxruntime/python/tools/transformers/onnx_model_bert.py
+++ b/onnxruntime/python/tools/transformers/onnx_model_bert.py
@@ -6,6 +6,7 @@
from logging import getLogger
from typing import List, Optional
+from convert_to_packing_mode import PackingMode
from fusion_attention import AttentionMask, FusionAttention
from fusion_biasgelu import FusionBiasGelu
from fusion_embedlayer import FusionEmbedLayerNormalization
@@ -482,3 +483,7 @@ class BertOnnxModel(OnnxModel):
logger.warning("Attention not fused")
return is_perfect
+
+ def convert_to_packing_mode(self, use_symbolic_shape_infer: bool = False):
+ packing_mode = PackingMode(self)
+ packing_mode.convert(use_symbolic_shape_infer)
diff --git a/onnxruntime/python/tools/transformers/optimizer.py b/onnxruntime/python/tools/transformers/optimizer.py
index 8614b18ee1..a3c16ebcae 100644
--- a/onnxruntime/python/tools/transformers/optimizer.py
+++ b/onnxruntime/python/tools/transformers/optimizer.py
@@ -410,6 +410,22 @@ def _parse_arguments():
)
parser.set_defaults(use_external_data_format=False)
+ parser.add_argument(
+ "--disable_symbolic_shape_infer",
+ required=False,
+ action="store_true",
+ help="diable symoblic shape inference",
+ )
+ parser.set_defaults(disable_symbolic_shape_infer=False)
+
+ parser.add_argument(
+ "--convert_to_packing_mode",
+ required=False,
+ action="store_true",
+ help="convert the model to packing mode. Only available for BERT like model",
+ )
+ parser.set_defaults(convert_to_packing_mode=False)
+
args = parser.parse_args()
return args
@@ -454,14 +470,20 @@ def main():
if args.input_int32:
optimizer.change_graph_inputs_to_int32()
- optimizer.save_model_to_file(args.output, args.use_external_data_format)
-
if args.model_type in ["bert", "gpt2"]:
if optimizer.is_fully_optimized():
logger.info("The model has been fully optimized.")
else:
logger.info("The model has been optimized.")
+ if args.convert_to_packing_mode:
+ if args.model_type == "bert":
+ optimizer.convert_to_packing_mode(not args.disable_symbolic_shape_infer)
+ else:
+ logger.warning("Packing mode only supports BERT like models")
+
+ optimizer.save_model_to_file(args.output, args.use_external_data_format)
+
if __name__ == "__main__":
main()
diff --git a/onnxruntime/test/contrib_ops/skiplayernorm_op_test.cc b/onnxruntime/test/contrib_ops/skiplayernorm_op_test.cc
index a6620eb528..6c0383b46c 100644
--- a/onnxruntime/test/contrib_ops/skiplayernorm_op_test.cc
+++ b/onnxruntime/test/contrib_ops/skiplayernorm_op_test.cc
@@ -24,14 +24,18 @@ static void RunTest(
int hidden_size,
bool use_float16 = false,
bool no_beta = false,
- bool simplified = false) {
+ bool simplified = false,
+ bool use_token_count = false) {
// Input and output shapes
- // Input 0 - input: (batch_size, sequence_length, hidden_size)
- // Input 1 - skip : (batch_size, sequence_length, hidden_size)
+ // Input 0 - input: (batch_size, sequence_length, hidden_size) or (batch_size * sequence_length, hidden_size)
+ // Input 1 - skip : (batch_size, sequence_length, hidden_size) or (batch_size * sequence_length, hidden_size)
// Input 2 - gamma: (hidden_size)
// Input 3 - beta : (hidden_size)
- // Output : (batch_size, sequence_length, hidden_size)
+ // Output : (batch_size, sequence_length, hidden_size) or (batch_size * sequence_length, hidden_size)
std::vector input_dims = {batch_size, sequence_length, hidden_size};
+ if (use_token_count) {
+ input_dims = {batch_size * sequence_length, hidden_size};
+ }
std::vector skip_dims = input_dims;
std::vector gamma_dims = {hidden_size};
std::vector beta_dims = gamma_dims;
@@ -504,6 +508,160 @@ TEST(SkipLayerNormTest, SkipLayerNormBatch2_Bias_ProducingOptionalOutput) {
hidden_size);
}
+TEST(SkipLayerNormTest, SkipLayerNormBatch1_Float16_vec_token_count) {
+ int batch_size = 1;
+ int sequence_length = 2;
+ int hidden_size = 64;
+
+ std::vector input_data = {
+ 0.8f, -0.5f, 0.0f, 1.f, 0.5f, 0.2f, 0.3f, -0.6f, // 1
+ -0.8f, -0.5f, 2.0f, 1.f, 0.5f, 0.2f, 0.3f, 0.2f, // 2
+ 0.8f, -0.5f, 0.0f, 1.f, -0.5f, 0.2f, 0.3f, 0.6f, // 3
+ 0.8f, -0.5f, 0.0f, 1.f, 0.5f, 0.1f, 0.3f, -0.3f, // 4
+ 0.8f, -3.5f, 0.9f, 1.f, 0.5f, 0.2f, 0.2f, -0.6f, // 5
+ 0.8f, -0.5f, 0.0f, 1.f, 0.5f, 0.2f, 0.3f, -0.6f, // 6
+ 0.9f, -0.5f, 0.8f, 2.f, 0.3f, 0.3f, 0.3f, -0.6f, // 7
+ 0.8f, -0.8f, 3.0f, 1.f, 0.5f, 0.2f, 0.3f, -0.6f, // 8
+ 0.8f, -0.5f, 0.1f, 1.f, 0.5f, 0.2f, 0.3f, -0.6f, // 9
+ 0.8f, -1.5f, 0.0f, 6.f, 0.5f, 0.2f, 0.3f, -0.6f, // 10
+ 0.8f, -0.5f, 0.0f, 2.f, 0.5f, 0.2f, 0.3f, -0.6f, // 11
+ 0.8f, -0.2f, 7.0f, 1.f, -0.2f, 0.2f, 0.3f, 0.6f, // 12
+ 0.8f, -0.5f, 0.0f, 1.f, 0.6f, 0.2f, 0.3f, -0.6f, // 13
+ 0.8f, -0.5f, 0.0f, 1.f, 0.5f, 0.3f, 0.3f, -0.6f, // 14
+ 0.8f, -0.5f, 0.0f, 1.f, 0.5f, 0.2f, -0.4f, 0.6f, // 15
+ 0.8f, -0.5f, 0.0f, 1.f, 0.5f, 0.2f, 0.3f, 0.1f}; // 16
+
+ std::vector skip_data = {
+ 0.8f, -0.5f, 0.0f, 1.f, 0.5f, 0.2f, 0.3f, -0.6f, // 1
+ -0.8f, -0.5f, 2.0f, 1.f, 0.5f, 0.2f, 0.3f, 0.2f, // 2
+ 0.8f, -0.5f, 0.0f, 1.f, -0.5f, 0.2f, 0.3f, 0.6f, // 3
+ 0.8f, -0.5f, 0.0f, 3.f, 0.5f, 0.1f, 0.3f, -0.4f, // 4
+ 0.8f, -3.5f, 2.9f, -0.f, 0.5f, 0.2f, 0.2f, 0.6f, // 5
+ 0.8f, -0.5f, 0.0f, 1.f, 0.5f, -0.2f, 0.3f, 0.6f, // 6
+ 0.9f, -0.5f, 0.8f, 2.f, 0.3f, 0.3f, 0.3f, -0.6f, // 7
+ 0.8f, -1.8f, 3.0f, 1.f, 0.5f, 0.2f, 0.3f, -0.6f, // 8
+ 0.8f, -0.5f, 0.1f, 1.f, 0.5f, 0.2f, 0.3f, -0.6f, // 9
+ 0.8f, -1.5f, 0.0f, 6.f, 0.5f, 0.2f, -1.2f, 0.6f, // 10
+ 0.8f, -3.5f, 0.0f, 2.f, -0.9f, 0.2f, 0.3f, 0.6f, // 11
+ 0.8f, -0.2f, 7.0f, 0.f, -0.2f, 0.2f, 0.3f, 0.6f, // 12
+ 0.8f, -0.5f, 4.0f, 1.f, 1.6f, 0.2f, 1.3f, -0.6f, // 13
+ 0.8f, -0.5f, 0.1f, 1.f, 0.5f, 0.3f, 0.3f, -0.6f, // 14
+ 0.8f, -0.5f, 1.0f, 0.f, 0.5f, 2.2f, -0.4f, 0.6f, // 15
+ 0.8f, -0.5f, 0.2f, 1.f, 0.5f, 0.2f, 0.3f, 0.1f}; // 16
+
+ std::vector gamma_data = {
+ 0.8f, -0.5f, 0.0f, 1.f, 0.5f, 0.2f, 4.3f, -0.6f, // 1
+ -0.8f, -3.5f, 2.0f, 1.f, 0.2f, 0.2f, 0.3f, 0.2f, // 2
+ 0.8f, -0.5f, 0.0f, 1.f, -0.5f, 0.2f, 0.3f, 0.6f, // 3
+ 0.8f, -0.5f, 0.0f, 1.f, 0.5f, 0.1f, 0.3f, -0.3f, // 4
+ 0.2f, -3.5f, 0.9f, -2.f, 0.5f, 1.2f, 0.2f, 0.6f, // 5
+ 0.8f, -0.5f, 0.0f, 1.f, 0.5f, 0.2f, 3.3f, -0.6f, // 6
+ 0.9f, -0.5f, -0.8f, 2.f, 0.3f, 0.3f, 0.3f, 0.6f, // 7
+ 0.1f, -0.5f, 0.0f, 1.f, 0.5f, 0.2f, 0.3f, 0.1f}; // 8
+
+ std::vector beta_data = {
+ 0.8f, -0.5f, 0.0f, 1.f, 0.5f, 0.2f, 0.3f, -0.6f, // 1
+ -0.8f, -0.5f, 2.0f, 0.f, 0.5f, 0.2f, 4.9f, 0.2f, // 2
+ 0.2f, -0.5f, 0.0f, 1.f, -0.5f, 0.2f, 0.3f, 0.6f, // 3
+ 0.8f, -0.5f, 0.0f, 1.f, 0.5f, 0.2f, 0.3f, -0.3f, // 4
+ 0.1f, -3.5f, 4.9f, 0.f, 0.5f, 0.2f, 0.2f, -0.6f, // 5
+ 0.8f, -1.5f, 0.0f, 3.f, 0.5f, 0.7f, 0.8f, -0.6f, // 6
+ 0.9f, -0.5f, 0.8f, 0.f, 0.3f, 0.3f, 0.3f, -0.6f, // 7
+ 0.8f, -0.5f, 0.0f, 1.f, 0.5f, 0.2f, 0.3f, 0.1f}; // 8
+
+ // Update test data result which use internal fp32 calculation for fp16 input/parameters.
+ // Following pytorch code snippet are used to generate the result: (Not torch uses fp32 internal calculation for this)
+ //
+ // gamma_tensor = torch.tensor(gamma_data, dtype=torch.float32).reshape(hidden_size).to('cuda:0').to(torch.float16)
+ // beta_tensor = torch.tensor(beta_data, dtype=torch.float32).reshape(hidden_size).to('cuda:0').to(torch.float16)
+ // input_tensor = torch.tensor(input_data, dtype=torch.float32).reshape(
+ // batch_size, sequence_length, hidden_size).to('cuda:0').to(torch.float16)
+ // skip_tensor = torch.tensor(skip_data, dtype=torch.float32).reshape(
+ // batch_size, sequence_length, hidden_size).to('cuda:0').to(torch.float16)
+ // added_input = torch.add(input_tensor, skip_tensor)
+ // out32 = torch.layer_norm(added_input, [hidden_size], gamma_tensor, beta_tensor, eps=epsilon_).to(torch.float32)
+ //
+ std::vector output_data = {
+ 1.25000000f, -0.04403687f, 0.00000000f, 1.79003906f, 0.61132812f, 0.17639160f, 0.28125000f, 0.01530457f,
+ 0.20166016f, 2.69140625f, 5.84765625f, 0.78955078f, 0.54443359f, 0.17639160f, 4.89843750f, 0.17639160f,
+ 0.64990234f, -0.04403687f, 0.00000000f, 1.79003906f, -0.04403687f, 0.17639160f, 0.29882812f, 0.80175781f,
+ 1.25000000f, -0.04403687f, 0.00000000f, 2.92382812f, 0.61132812f, 0.17687988f, 0.29882812f, -0.07745361f,
+ 0.21240234f, 11.60156250f, 6.52734375f, -0.44482422f, 0.61132812f, 0.05844116f, 0.17639160f, -0.80712891f,
+ 1.25000000f, -1.04394531f, 0.00000000f, 3.78906250f, 0.61132812f, 0.63134766f, 0.78564453f, -0.39331055f,
+ 1.50878906f, -0.04403687f, 0.34985352f, 3.84765625f, 0.29882812f, 0.29882812f, 0.29882812f, -1.21582031f,
+ 0.85595703f, 0.40966797f, 0.00000000f, 1.79003906f, 0.61132812f, 0.17639160f, 0.29882812f, -0.00255013f,
+ 1.00976562f, -0.12152100f, 0.00000000f, 1.41894531f, 0.51367188f, 0.15832520f, -0.25805664f, -0.09875488f,
+ -1.00976562f, 4.89453125f, 1.26953125f, 4.33984375f, 0.50537109f, 0.15832520f, 4.68359375f, 0.12695312f,
+ 0.40942383f, 0.46655273f, 0.00000000f, 2.20312500f, -0.23913574f, 0.15832520f, 0.26123047f, 0.38110352f,
+ 1.00976562f, -0.23913574f, 0.00000000f, 1.02734375f, 0.23913574f, 0.17907715f, 0.26123047f, -0.33178711f,
+ 0.15234375f, -0.85058594f, 5.98046875f, -0.83789062f, 0.74853516f, -0.04998779f, 0.25244141f, -1.10156250f,
+ 1.00976562f, -1.12109375f, 0.00000000f, 3.41992188f, 0.51367188f, 0.67431641f, 0.37133789f, -0.09875488f,
+ 1.13574219f, -0.12152100f, 0.77832031f, 0.05398560f, 0.30810547f, 0.47265625f, 0.09643555f, -0.53662109f,
+ 0.82617188f, -0.12152100f, 0.00000000f, 1.41894531f, 0.51367188f, 0.15832520f, 0.26123047f, 0.07135010f};
+
+ RunTest(input_data,
+ skip_data,
+ gamma_data,
+ beta_data,
+ std::vector(),
+ output_data,
+ {},
+ epsilon_,
+ batch_size,
+ sequence_length,
+ hidden_size,
+ true,
+ false,
+ false,
+ true);
+}
+
+TEST(SkipLayerNormTest, SkipLayerNormBatch2_TokenCount) {
+ int batch_size = 2;
+ int sequence_length = 2;
+ int hidden_size = 4;
+
+ std::vector input_data = {
+ 0.8f, -0.5f, 0.0f, 1.f,
+ 0.5f, 0.2f, 0.3f, -0.6f,
+ 0.8f, -0.5f, 0.0f, 1.f,
+ 0.5f, 0.2f, 0.3f, -0.6f};
+
+ std::vector skip_data = {
+ 0.1f, -0.2f, 0.3f, 1.0f,
+ 0.5f, 0.1f, 0.4f, 1.6f,
+ 1.8f, -0.3f, 0.0f, 1.f,
+ -0.5f, 0.4f, 0.8f, -0.6f};
+
+ std::vector gamma_data = {
+ 0.3f, 0.2f, 4.0f, 2.2f};
+
+ std::vector beta_data = {
+ 0.2f, 0.1f, 0.4f, 1.6f};
+
+ std::vector output_data = {
+ 0.28433859348297119, -0.17090578377246857, -0.92897164821624756, 4.6924152374267578,
+ 0.46111652255058289, -0.21333980560302734, -0.29631003737449646, 3.5148544311523438,
+ 0.55470430850982666, -0.15080101788043976, -2.3229825496673584, 3.255286693572998,
+ 0.15631480515003204, 0.21066918969154358, 4.9432611465454102, -1.7957965135574341};
+
+ RunTest(input_data,
+ skip_data,
+ gamma_data,
+ beta_data,
+ std::vector(),
+ output_data,
+ {},
+ epsilon_,
+ batch_size,
+ sequence_length,
+ hidden_size,
+ false,
+ false,
+ false,
+ true);
+}
+
// SkipSimplifiedLayerNorm has not been enabled for ROCm and DML yet
#if !defined(USE_ROCM) && !defined(USE_DML)
TEST(SkipLayerNormTest, SkipSimplifiedLayerNormBatch1_Float16) {