cleanup formatting in skip_layer_norm.cc (#8371)

This commit is contained in:
Nick Kreeger 2021-07-13 16:36:41 -05:00 committed by GitHub
parent 31f291f0af
commit 178c139718
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23

View file

@ -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();
}