From de77c60e6e81212937d4c6e6d2b1ae8830a86ee0 Mon Sep 17 00:00:00 2001 From: Sunny Shukla <102830081+sunnyshu-intel@users.noreply.github.com> Date: Wed, 16 Nov 2022 21:25:00 -0800 Subject: [PATCH] [oneDNN ep] SLN performance improvement for bias (#13620) ### Description SkipLayerNorm performance improvement when bias is present as input ### Motivation and Context - For SkipLayerNorm op, adding bias tensor using post-op to the add primitive adding input and skip tensors is causing drastic performance degradation. - Hence the post-op is removed and instead, two add primitives are used in series, adding input and skip, and then adding bias to the result of input and skip. - This change has shown a significant amount of performance gain for SkipLayerNorm operator. --- .../providers/dnnl/subgraph/dnnl_layernorm.cc | 150 ++++++++---------- 1 file changed, 69 insertions(+), 81 deletions(-) diff --git a/onnxruntime/core/providers/dnnl/subgraph/dnnl_layernorm.cc b/onnxruntime/core/providers/dnnl/subgraph/dnnl_layernorm.cc index 66ace706a5..7d3d26bc97 100644 --- a/onnxruntime/core/providers/dnnl/subgraph/dnnl_layernorm.cc +++ b/onnxruntime/core/providers/dnnl/subgraph/dnnl_layernorm.cc @@ -44,16 +44,18 @@ Skip Layer Norm: 1) M - Mean Tensor (Optional) 2) I - Inverse std Tensor (Optional) - +-----------+ -(X) ---------->+ | +-----------+ - | | (X + S + B) | | -(S) ---------->+ BuildSLN +------------------>+ +----------> (Y) - | | | | -(B) ---------->+ | (G) -------->+ LayerNorm +----------> (M) - +-----------+ | | - (E) -------->+ +----------> (I) - | | - +-----------+ +(X) ---------->+-----------+ + | Add | +(S) ---------->+-----------+ + | +-----------+ + | (X + S + B) | | + +-----------+-------------->+ +----------> (Y) + | Add | | | +(B) ---------------->+-----------+ (G) ------->+ LayerNorm +----------> (M) + | | + (E) ------->+ +----------> (I) + | | + +-----------+ Attributes (epsilon) */ @@ -71,9 +73,11 @@ void DnnlLayerNorm::CreatePrimitive(DnnlSubgraphPrimitive& sp, DnnlNode& node) { // Input positions int shift_pos, scale_pos; - // Get src mem - auto src_mem = sp.GetMemory(node.Input(IN_INPUT)); - auto src_md = src_mem.get_desc(); + // Get src desc + auto src_md = sp.GetMemory(node.Input(IN_INPUT)).get_desc(); + + // Init src mem + dnnl::memory src_mem; // This contains the layer norm op and its parameters ln_components op_comps; @@ -85,9 +89,57 @@ void DnnlLayerNorm::CreatePrimitive(DnnlSubgraphPrimitive& sp, DnnlNode& node) { // Fix positions for arguments shift_pos = IN_BETA; scale_pos = IN_SLN_GAMMA; - - // Build SLN and get modified mem - src_mem = BuildSLN(sp, node, dnnl_engine); + + // Move the src to GPU if needed + src_mem = sp.GetMemoryAndReshape(node.Input(IN_INPUT), src_md, dnnl_engine); + + // Make dst desc, must be same as src + auto dst_md = dnnl::memory::desc(src_md.dims(), node.Output(OUT_OUTPUT).Type(), dnnl::memory::format_tag::any); + + // Add src + skip + { + + // get skip desc + auto skip_md = sp.GetMemory(node.Input(IN_SKIP)).get_desc(); + // Move the skip to GPU if needed + auto skip_mem = sp.GetMemoryAndReshape(node.Input(IN_SKIP), skip_md, dnnl_engine); + + // Create and add primitive + auto add_skip_d = dnnl::binary::desc(dnnl::algorithm::binary_add, src_md, skip_md, dst_md); + auto add_skip_pd = dnnl::binary::primitive_desc(add_skip_d, dnnl_engine); + auto add_skip = dnnl::binary(add_skip_pd); + std::unordered_map add_skip_mem_map({{DNNL_ARG_SRC_0, src_mem}, {DNNL_ARG_SRC_1, skip_mem}, {DNNL_ARG_DST, src_mem}}); + sp.AddPrimitive(add_skip, add_skip_mem_map); + + } + + // Add src + skip + bias + if (node.Input(IN_SLN_BIAS).Exists()) { + + // get bias desc + auto bias_md = sp.GetMemory(node.Input(IN_SLN_BIAS)).get_desc(); + // Move the bias to GPU if needed + auto bias_mem = sp.GetMemoryAndReshape(node.Input(IN_SLN_BIAS), bias_md, dnnl_engine); + // Get bias dims + auto bias_dims = bias_md.dims(); + // Get src dims + auto src_dims = src_md.dims(); + + // To follow the spec means our bias will always have less dimensions that our input + // so we add the extra dimensions, reshape it and let OneDNN broadcast the value + while (bias_dims.size() < src_dims.size()) { + bias_dims.insert(bias_dims.begin(), 1); + } + bias_md = bias_md.reshape(bias_dims); + + // Create and add primitive + auto add_bias_d = dnnl::binary::desc(dnnl::algorithm::binary_add, src_md, bias_md, dst_md); + auto add_bias_pd = dnnl::binary::primitive_desc(add_bias_d, dnnl_engine); + auto add_bias = dnnl::binary(add_bias_pd); + std::unordered_map add_bias_mem_map({{DNNL_ARG_SRC_0, src_mem}, {DNNL_ARG_SRC_1, bias_mem}, {DNNL_ARG_DST, src_mem}}); + sp.AddPrimitive(add_bias, add_bias_mem_map); + + } } else if (node.OpType() == "LayerNormalization") { @@ -99,7 +151,7 @@ void DnnlLayerNorm::CreatePrimitive(DnnlSubgraphPrimitive& sp, DnnlNode& node) { scale_pos = IN_LN_GAMMA; // Move the src to GPU if needed - src_mem = sp.GetMemoryAndReshape(node.Input(IN_INPUT), src_mem.get_desc(), dnnl_engine); + src_mem = sp.GetMemoryAndReshape(node.Input(IN_INPUT), src_md, dnnl_engine); } else { ORT_THROW("Unknown LayerNormalization flavor"); @@ -195,70 +247,6 @@ void DnnlLayerNorm::CreatePrimitive(DnnlSubgraphPrimitive& sp, DnnlNode& node) { sp.SetMemory(node.Output(OUT_OUTPUT), src_mem, true); } -dnnl::memory DnnlLayerNorm::BuildSLN(DnnlSubgraphPrimitive& sp, DnnlNode& node, dnnl::engine dnnl_engine) { - // X += SKIP - // Get input and skip info - auto input_md = sp.GetMemory(node.Input(IN_INPUT)).get_desc(); - auto skip_md = sp.GetMemory(node.Input(IN_SKIP)).get_desc(); - auto skip_dims = skip_md.dims(); - - // Create primitive input map - std::unordered_map skip_bias_args; - - // Create md for the add op, according to the spec the output type should be - // the same as the input, so we can support inplace ops - auto add_skip_dst_md = dnnl::memory::desc(skip_dims, node.Output(OUT_OUTPUT).Type(), dnnl::memory::format_tag::any); - // Create desc for the op and primitive - auto add_skip_d = dnnl::binary::desc(dnnl::algorithm::binary_add, input_md, skip_md, add_skip_dst_md); - - // Create primitive descriptor container - dnnl::binary::primitive_desc add_skip_pd; - // Add post op bias - if (node.Input(IN_SLN_BIAS).Exists()) { - // X += BIAS - // Get bias md - auto bias_md = sp.GetMemory(node.Input(IN_SLN_BIAS)).get_desc(); - auto bias_dims = bias_md.dims(); - // To follow the spec means our bias will always have less dimensions that our input} - // so we add the extra dimensions, reshape it and let OneDNN broadcast the value - while (bias_dims.size() < skip_dims.size()) { - bias_dims.insert(bias_dims.begin(), 1); - } - bias_md = bias_md.reshape(bias_dims); - - dnnl::post_ops bias_add; - dnnl::primitive_attr binary_attr; - bias_add.append_binary(dnnl::algorithm::binary_add, bias_md); - binary_attr.set_post_ops(bias_add); - // Add post op to scale result - add_skip_pd = dnnl::binary::primitive_desc(add_skip_d, binary_attr, dnnl_engine); - - // Get bias mem - auto bias_mem = sp.GetMemoryAndReshape(node.Input(IN_SLN_BIAS), bias_md, dnnl_engine); - // Add bias arg - skip_bias_args.insert({DNNL_ARG_ATTR_MULTIPLE_POST_OP(0) | DNNL_ARG_SRC_1, bias_mem}); - - } else { - add_skip_pd = dnnl::binary::primitive_desc(add_skip_d, dnnl_engine); - } - - // Move the memory to the target device - auto src_mem = sp.GetMemoryAndReshape(node.Input(IN_INPUT), add_skip_pd.src0_desc(), dnnl_engine); - auto skip_mem = sp.GetMemoryAndReshape(node.Input(IN_SKIP), add_skip_pd.src1_desc(), dnnl_engine); - - // Add args - skip_bias_args.insert({DNNL_ARG_SRC_0, src_mem}); - skip_bias_args.insert({DNNL_ARG_SRC_1, skip_mem}); - skip_bias_args.insert({DNNL_ARG_DST, src_mem}); - - // Create and add primitive - auto add_skip_prim = dnnl::binary(add_skip_pd); - sp.AddPrimitive(add_skip_prim, skip_bias_args); - - // Return src - return src_mem; -} - void DnnlLayerNorm::ValidateDims(DnnlSubgraphPrimitive& sp, DnnlNode& node) { // Get input and evaluate auto input_dims = sp.GetMemory(node.Input(IN_INPUT)).get_desc().dims();