[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.
This commit is contained in:
Sunny Shukla 2022-11-16 21:25:00 -08:00 committed by GitHub
parent b731cf397d
commit de77c60e6e
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23

View file

@ -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<int, dnnl::memory> 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<int, dnnl::memory> 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<int, dnnl::memory> 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();