diff --git a/onnxruntime/contrib_ops/cpu/skip_layer_norm.cc b/onnxruntime/contrib_ops/cpu/skip_layer_norm.cc index b84635d082..6119bc20ee 100644 --- a/onnxruntime/contrib_ops/cpu/skip_layer_norm.cc +++ b/onnxruntime/contrib_ops/cpu/skip_layer_norm.cc @@ -98,36 +98,38 @@ Status SkipLayerNorm::Compute(OpKernelContext* p_ctx) const { T* output_data = output->MutableData(); - concurrency::ThreadPool::TryBatchParallelFor(p_ctx->GetOperatorThreadPool(), static_cast(task_count), - [&](ptrdiff_t task_idx) { - const T* p_input = input_data + task_idx * hidden_size; - const T* p_skip = skip_data + task_idx * hidden_size; - T* p_output = output_data + task_idx * hidden_size; + concurrency::ThreadPool::TryBatchParallelFor( + p_ctx->GetOperatorThreadPool(), static_cast(task_count), + [&](ptrdiff_t task_idx) { + const T* p_input = input_data + task_idx * hidden_size; + const T* p_skip = skip_data + task_idx * hidden_size; + T* p_output = output_data + task_idx * hidden_size; - T mean = 0; - T mean_square = 0; + T mean = 0; + T mean_square = 0; - for (int64_t h = 0; h < hidden_size; h++) { - T value = p_input[h] + p_skip[h]; - if (nullptr != bias_data) { - value += bias_data[h]; - } - p_output[h] = value; - mean += value; - mean_square += value * value; - } + for (int64_t h = 0; h < hidden_size; h++) { + T value = p_input[h] + p_skip[h]; + if (nullptr != bias_data) { + value += bias_data[h]; + } + p_output[h] = value; + mean += value; + mean_square += value * value; + } - mean = mean / hidden_size; - mean_square = sqrt(mean_square / hidden_size - mean * mean + epsilon_); + mean = mean / hidden_size; + mean_square = sqrt(mean_square / hidden_size - mean * mean + epsilon_); - for (int64_t h = 0; h < hidden_size; h++) { - if (nullptr == beta_data) { - p_output[h] = (p_output[h] - mean) / mean_square * gamma_data[h]; - } else { - p_output[h] = (p_output[h] - mean) / mean_square * gamma_data[h] + beta_data[h]; - } - } - }, 0); + for (int64_t h = 0; h < hidden_size; h++) { + if (nullptr == beta_data) { + p_output[h] = (p_output[h] - mean) / mean_square * gamma_data[h]; + } else { + p_output[h] = (p_output[h] - mean) / mean_square * gamma_data[h] + beta_data[h]; + } + } + }, + 0); return Status::OK(); }