mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
cleanup formatting in skip_layer_norm.cc (#8371)
This commit is contained in:
parent
31f291f0af
commit
178c139718
1 changed files with 28 additions and 26 deletions
|
|
@ -98,36 +98,38 @@ Status SkipLayerNorm<T>::Compute(OpKernelContext* p_ctx) const {
|
|||
|
||||
T* output_data = output->MutableData<T>();
|
||||
|
||||
concurrency::ThreadPool::TryBatchParallelFor(p_ctx->GetOperatorThreadPool(), static_cast<int32_t>(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<int32_t>(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();
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in a new issue