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) {