From 46888e4744792743beea95516e31b84890f66370 Mon Sep 17 00:00:00 2001 From: Reza Yazdani Date: Tue, 8 Sep 2020 21:29:37 +0000 Subject: [PATCH 1/8] fixing the compilation error for AVX2 architecture --- csrc/adam/cpu_adam.cpp | 783 +++++++++++++++++++-------------------- csrc/includes/cpu_adam.h | 40 +- setup.py | 14 +- 3 files changed, 435 insertions(+), 402 deletions(-) mode change 100644 => 100755 csrc/adam/cpu_adam.cpp mode change 100644 => 100755 csrc/includes/cpu_adam.h diff --git a/csrc/adam/cpu_adam.cpp b/csrc/adam/cpu_adam.cpp old mode 100644 new mode 100755 index 2baab6177c2e..bc0c32e45f52 --- a/csrc/adam/cpu_adam.cpp +++ b/csrc/adam/cpu_adam.cpp @@ -25,35 +25,42 @@ void Adam_Optimizer::Step(float* _params, size_t _param_size, __half* dev_params) { - _betta1_t *= _betta1; - _betta2_t *= _betta2; - - AVX_512 betta1_4; - betta1_4.data = _mm512_set1_ps(_betta1); - AVX_512 betta2_4; - betta2_4.data = _mm512_set1_ps(_betta2); - float betta1_minus1 = 1 - _betta1; float betta2_minus1 = 1 - _betta2; - AVX_512 betta1_minus1_4; - betta1_minus1_4.data = _mm512_set1_ps(betta1_minus1); - AVX_512 betta2_minus1_4; - betta2_minus1_4.data = _mm512_set1_ps(betta2_minus1); float bias_correction1 = 1 - _betta1_t; float bias_correction2 = 1 / sqrt(1 - _betta2_t); - // AVX_512 bias_correction1_4 = _mm512_set1_ps(bias_correction1); - AVX_512 bias2_sqrt; - bias2_sqrt.data = _mm512_set1_ps(bias_correction2); - - AVX_512 eps_4; - eps_4.data = _mm512_set1_ps(_eps); float step_size = -1 * _alpha / bias_correction1; - AVX_512 step_size_4; - step_size_4.data = _mm512_set1_ps(step_size); - size_t rounded_size = ROUND_DOWN(_param_size, SIMD_WIDTH); + size_t rounded_size = 0; + +#if defined(__AVX512__) or defined(__AVX256__) + + AVX_Data betta1_4; + betta1_4.data = SIMD_SET(_betta1); + AVX_Data betta2_4; + betta2_4.data = SIMD_SET(_betta2); + + AVX_Data betta1_minus1_4; + betta1_minus1_4.data = SIMD_SET(betta1_minus1); + AVX_Data betta2_minus1_4; + betta2_minus1_4.data = SIMD_SET(betta2_minus1); + + AVX_Data bias2_sqrt; + bias2_sqrt.data = SIMD_SET(bias_correction2); + + AVX_Data eps_4; + eps_4.data = SIMD_SET(_eps); + + AVX_Data step_size_4; + step_size_4.data = SIMD_SET(step_size); + + AVX_Data weight_decay4; + if (_weight_decay > 0) + weight_decay4.data = SIMD_SET(_weight_decay); + + rounded_size = ROUND_DOWN(_param_size, SIMD_WIDTH); for (size_t t = 0; t < rounded_size; t += TILE) { size_t copy_size = TILE; @@ -61,57 +68,41 @@ void Adam_Optimizer::Step(float* _params, size_t offset = copy_size + t; #pragma omp parallel for for (size_t i = t; i < offset; i += SIMD_WIDTH) { - AVX_512 grad_4; - grad_4.data = _mm512_loadu_ps(grads + i); + AVX_Data grad_4; + grad_4.data = SIMD_LOAD(grads + i); - AVX_512 momntum_4; - momntum_4.data = _mm512_loadu_ps(_exp_avg + i); - AVX_512 varianc_4; - varianc_4.data = _mm512_loadu_ps(_exp_avg_sq + i); + AVX_Data momentum_4; + momentum_4.data = SIMD_LOAD(_exp_avg + i); + AVX_Data variance_4; + variance_4.data = SIMD_LOAD(_exp_avg_sq + i); - AVX_512 param_4; - param_4.data = _mm512_loadu_ps(_params + i); + AVX_Data param_4; + param_4.data = SIMD_LOAD(_params + i); - if (_weight_decay > 0) { - AVX_512 weight_decay4; - weight_decay4.data = _mm512_set1_ps(_weight_decay); - grad_4.data = _mm512_fmadd_ps(param_4.data, weight_decay4.data, grad_4.data); - } + if (_weight_decay > 0) + grad_4.data = SIMD_FMA(param_4.data, weight_decay4.data, grad_4.data); - momntum_4.data = _mm512_mul_ps(momntum_4.data, betta1_4.data); - momntum_4.data = _mm512_fmadd_ps(grad_4.data, betta1_minus1_4.data, momntum_4.data); + momentum_4.data = SIMD_MUL(momentum_4.data, betta1_4.data); + momentum_4.data = SIMD_FMA(grad_4.data, betta1_minus1_4.data, momentum_4.data); - varianc_4.data = _mm512_mul_ps(varianc_4.data, betta2_4.data); - grad_4.data = _mm512_mul_ps(grad_4.data, grad_4.data); - varianc_4.data = _mm512_fmadd_ps(grad_4.data, betta2_minus1_4.data, varianc_4.data); + variance_4.data = SIMD_MUL(variance_4.data, betta2_4.data); + grad_4.data = SIMD_MUL(grad_4.data, grad_4.data); + variance_4.data = SIMD_FMA(grad_4.data, betta2_minus1_4.data, variance_4.data); - grad_4.data = _mm512_sqrt_ps(varianc_4.data); - grad_4.data = _mm512_fmadd_ps(grad_4.data, bias2_sqrt.data, eps_4.data); - grad_4.data = _mm512_div_ps(momntum_4.data, grad_4.data); + grad_4.data = SIMD_SQRT(variance_4.data); + grad_4.data = SIMD_FMA(grad_4.data, bias2_sqrt.data, eps_4.data); + grad_4.data = SIMD_DIV(momentum_4.data, grad_4.data); - param_4.data = _mm512_fmadd_ps(grad_4.data, step_size_4.data, param_4.data); + param_4.data = SIMD_FMA(grad_4.data, step_size_4.data, param_4.data); - _mm512_storeu_ps(_params + i, param_4.data); + SIMD_STORE(_params + i, param_4.data); - if (dev_params) _mm512_storeu_ps(_doubled_buffer[_buf_index] + (i - t), param_4.data); + if (dev_params) SIMD_STORE(_doubled_buffer[_buf_index] + (i - t), param_4.data); - _mm512_storeu_ps(_exp_avg + i, momntum_4.data); - _mm512_storeu_ps(_exp_avg_sq + i, varianc_4.data); + SIMD_STORE(_exp_avg + i, momentum_4.data); + SIMD_STORE(_exp_avg_sq + i, variance_4.data); } - if (dev_params) { /* - #pragma omp parallel for - for (size_t j = 0; j < copy_size; j += 4) { - _doubled_buffer[_buf_index][j] = (__half)_params[t + j]; - _doubled_buffer[_buf_index][j + 1] = (__half)_params[t + j + 1]; - _doubled_buffer[_buf_index][j + 2] = (__half)_params[t + j + 2]; - _doubled_buffer[_buf_index][j + 3] = (__half)_params[t + j + 3]; - } - - CUDA_CHECK(cudaMemcpyAsync(dev_params + t, - _doubled_buffer[_buf_index], - copy_size * sizeof(__half), - cudaMemcpyHostToDevice, - Context::Instance().GetCurrentStream()));*/ + if (dev_params) { launch_param_update(_doubled_buffer[_buf_index], dev_params + t, copy_size, @@ -120,32 +111,34 @@ void Adam_Optimizer::Step(float* _params, } } +#endif + if (_param_size > rounded_size) { #pragma omp parallel for for (size_t k = rounded_size; k < _param_size; k++) { float grad = grads[k]; float param = _params[k]; - float momntum = _exp_avg[k]; - float varianc = _exp_avg_sq[k]; - if (_weight_decay > 0) { grad = param * _weight_decay + grad; } + float momentum = _exp_avg[k]; + float variance = _exp_avg_sq[k]; + if (_weight_decay > 0) grad = param * _weight_decay + grad; - momntum *= momntum * _betta1; - momntum = grad * betta1_minus1 + momntum; + momentum *= momentum * _betta1; + momentum = grad * betta1_minus1 + momentum; - varianc = varianc * _betta2; + variance = variance * _betta2; grad = grad * grad; - varianc = grad * betta2_minus1 + varianc; + variance = grad * betta2_minus1 + variance; - grad = sqrt(varianc); + grad = sqrt(variance); grad = grad * bias_correction2 + _eps; - grad = momntum / grad; + grad = momentum / grad; param = grad * step_size + param; if (dev_params) _doubled_buffer[_buf_index][k - rounded_size] = (__half)param; _params[k] = param; - _exp_avg[k] = momntum; - _exp_avg_sq[k] = varianc; + _exp_avg[k] = momentum; + _exp_avg_sq[k] = variance; } if (dev_params) { launch_param_update(_doubled_buffer[_buf_index], @@ -163,35 +156,36 @@ void Adam_Optimizer::Step_4(float* _params, size_t _param_size, __half* dev_params) { - _betta1_t *= _betta1; - _betta2_t *= _betta2; + size_t rounded_size = 0; + +#if defined(__AVX512__) or defined(__AVX256__) - AVX_512 betta1_4; - betta1_4.data = _mm512_set1_ps(_betta1); - AVX_512 betta2_4; - betta2_4.data = _mm512_set1_ps(_betta2); + AVX_Data betta1_4; + betta1_4.data = SIMD_SET(_betta1); + AVX_Data betta2_4; + betta2_4.data = SIMD_SET(_betta2); float betta1_minus1 = 1 - _betta1; float betta2_minus1 = 1 - _betta2; - AVX_512 betta1_minus1_4; - betta1_minus1_4.data = _mm512_set1_ps(betta1_minus1); - AVX_512 betta2_minus1_4; - betta2_minus1_4.data = _mm512_set1_ps(betta2_minus1); + AVX_Data betta1_minus1_4; + betta1_minus1_4.data = SIMD_SET(betta1_minus1); + AVX_Data betta2_minus1_4; + betta2_minus1_4.data = SIMD_SET(betta2_minus1); float bias_correction1 = 1 - _betta1_t; float bias_correction2 = 1 / sqrt(1 - _betta2_t); - // AVX_512 bias_correction1_4 = _mm512_set1_ps(bias_correction1); - AVX_512 bias2_sqrt; - bias2_sqrt.data = _mm512_set1_ps(bias_correction2); + // AVX_Data bias_correction1_4 = SIMD_SET(bias_correction1); + AVX_Data bias2_sqrt; + bias2_sqrt.data = SIMD_SET(bias_correction2); - AVX_512 eps_4; - eps_4.data = _mm512_set1_ps(_eps); + AVX_Data eps_4; + eps_4.data = SIMD_SET(_eps); float step_size = -1 * _alpha / bias_correction1; - AVX_512 step_size_4; - step_size_4.data = _mm512_set1_ps(step_size); + AVX_Data step_size_4; + step_size_4.data = SIMD_SET(step_size); - size_t rounded_size = ROUND_DOWN(_param_size, (SIMD_WIDTH << 2)); + rounded_size = ROUND_DOWN(_param_size, (SIMD_WIDTH << 2)); for (size_t t = 0; t < rounded_size; t += TILE) { size_t copy_size = TILE; @@ -199,133 +193,119 @@ void Adam_Optimizer::Step_4(float* _params, size_t offset = copy_size + t; #pragma omp parallel for for (size_t i = t; i < offset; i += (SIMD_WIDTH << 2)) { - AVX_512 grad_4[4]; - grad_4[0].data = _mm512_loadu_ps(grads + i); - grad_4[1].data = _mm512_loadu_ps(grads + i + SIMD_WIDTH); - grad_4[2].data = _mm512_loadu_ps(grads + i + (SIMD_WIDTH << 1)); - grad_4[3].data = _mm512_loadu_ps(grads + i + SIMD_WIDTH * 3); - - AVX_512 momntum_4[4]; - momntum_4[0].data = _mm512_loadu_ps(_exp_avg + i); - momntum_4[1].data = _mm512_loadu_ps(_exp_avg + i + SIMD_WIDTH); - momntum_4[2].data = _mm512_loadu_ps(_exp_avg + i + (SIMD_WIDTH << 1)); - momntum_4[3].data = _mm512_loadu_ps(_exp_avg + i + SIMD_WIDTH * 3); - - AVX_512 varianc_4[4]; - varianc_4[0].data = _mm512_loadu_ps(_exp_avg_sq + i); - varianc_4[1].data = _mm512_loadu_ps(_exp_avg_sq + i + SIMD_WIDTH); - varianc_4[2].data = _mm512_loadu_ps(_exp_avg_sq + i + (SIMD_WIDTH << 1)); - varianc_4[3].data = _mm512_loadu_ps(_exp_avg_sq + i + SIMD_WIDTH * 3); - - AVX_512 param_4[4]; - param_4[0].data = _mm512_loadu_ps(_params + i); - param_4[1].data = _mm512_loadu_ps(_params + i + SIMD_WIDTH); - param_4[2].data = _mm512_loadu_ps(_params + i + (SIMD_WIDTH << 1)); - param_4[3].data = _mm512_loadu_ps(_params + i + SIMD_WIDTH * 3); + AVX_Data grad_4[4]; + grad_4[0].data = SIMD_LOAD(grads + i); + grad_4[1].data = SIMD_LOAD(grads + i + SIMD_WIDTH); + grad_4[2].data = SIMD_LOAD(grads + i + (SIMD_WIDTH << 1)); + grad_4[3].data = SIMD_LOAD(grads + i + SIMD_WIDTH * 3); + + AVX_Data momentum_4[4]; + momentum_4[0].data = SIMD_LOAD(_exp_avg + i); + momentum_4[1].data = SIMD_LOAD(_exp_avg + i + SIMD_WIDTH); + momentum_4[2].data = SIMD_LOAD(_exp_avg + i + (SIMD_WIDTH << 1)); + momentum_4[3].data = SIMD_LOAD(_exp_avg + i + SIMD_WIDTH * 3); + + AVX_Data variance_4[4]; + variance_4[0].data = SIMD_LOAD(_exp_avg_sq + i); + variance_4[1].data = SIMD_LOAD(_exp_avg_sq + i + SIMD_WIDTH); + variance_4[2].data = SIMD_LOAD(_exp_avg_sq + i + (SIMD_WIDTH << 1)); + variance_4[3].data = SIMD_LOAD(_exp_avg_sq + i + SIMD_WIDTH * 3); + + AVX_Data param_4[4]; + param_4[0].data = SIMD_LOAD(_params + i); + param_4[1].data = SIMD_LOAD(_params + i + SIMD_WIDTH); + param_4[2].data = SIMD_LOAD(_params + i + (SIMD_WIDTH << 1)); + param_4[3].data = SIMD_LOAD(_params + i + SIMD_WIDTH * 3); if (_weight_decay > 0) { - AVX_512 weight_decay4; - weight_decay4.data = _mm512_set1_ps(_weight_decay); + AVX_Data weight_decay4; + weight_decay4.data = SIMD_SET(_weight_decay); grad_4[0].data = - _mm512_fmadd_ps(param_4[0].data, weight_decay4.data, grad_4[0].data); + SIMD_FMA(param_4[0].data, weight_decay4.data, grad_4[0].data); grad_4[1].data = - _mm512_fmadd_ps(param_4[1].data, weight_decay4.data, grad_4[1].data); + SIMD_FMA(param_4[1].data, weight_decay4.data, grad_4[1].data); grad_4[2].data = - _mm512_fmadd_ps(param_4[2].data, weight_decay4.data, grad_4[2].data); + SIMD_FMA(param_4[2].data, weight_decay4.data, grad_4[2].data); grad_4[3].data = - _mm512_fmadd_ps(param_4[3].data, weight_decay4.data, grad_4[3].data); + SIMD_FMA(param_4[3].data, weight_decay4.data, grad_4[3].data); } - momntum_4[0].data = _mm512_mul_ps(momntum_4[0].data, betta1_4.data); - momntum_4[0].data = - _mm512_fmadd_ps(grad_4[0].data, betta1_minus1_4.data, momntum_4[0].data); - momntum_4[1].data = _mm512_mul_ps(momntum_4[1].data, betta1_4.data); - momntum_4[1].data = - _mm512_fmadd_ps(grad_4[1].data, betta1_minus1_4.data, momntum_4[1].data); - momntum_4[2].data = _mm512_mul_ps(momntum_4[2].data, betta1_4.data); - momntum_4[2].data = - _mm512_fmadd_ps(grad_4[2].data, betta1_minus1_4.data, momntum_4[2].data); - momntum_4[3].data = _mm512_mul_ps(momntum_4[3].data, betta1_4.data); - momntum_4[3].data = - _mm512_fmadd_ps(grad_4[3].data, betta1_minus1_4.data, momntum_4[3].data); - - varianc_4[0].data = _mm512_mul_ps(varianc_4[0].data, betta2_4.data); - varianc_4[1].data = _mm512_mul_ps(varianc_4[1].data, betta2_4.data); - varianc_4[2].data = _mm512_mul_ps(varianc_4[2].data, betta2_4.data); - varianc_4[3].data = _mm512_mul_ps(varianc_4[3].data, betta2_4.data); - grad_4[0].data = _mm512_mul_ps(grad_4[0].data, grad_4[0].data); - grad_4[1].data = _mm512_mul_ps(grad_4[1].data, grad_4[1].data); - grad_4[2].data = _mm512_mul_ps(grad_4[2].data, grad_4[2].data); - grad_4[3].data = _mm512_mul_ps(grad_4[3].data, grad_4[3].data); - varianc_4[0].data = - _mm512_fmadd_ps(grad_4[0].data, betta2_minus1_4.data, varianc_4[0].data); - varianc_4[1].data = - _mm512_fmadd_ps(grad_4[1].data, betta2_minus1_4.data, varianc_4[1].data); - varianc_4[2].data = - _mm512_fmadd_ps(grad_4[2].data, betta2_minus1_4.data, varianc_4[2].data); - varianc_4[3].data = - _mm512_fmadd_ps(grad_4[3].data, betta2_minus1_4.data, varianc_4[3].data); - - grad_4[0].data = _mm512_sqrt_ps(varianc_4[0].data); - grad_4[1].data = _mm512_sqrt_ps(varianc_4[1].data); - grad_4[2].data = _mm512_sqrt_ps(varianc_4[2].data); - grad_4[3].data = _mm512_sqrt_ps(varianc_4[3].data); - - grad_4[0].data = _mm512_fmadd_ps(grad_4[0].data, bias2_sqrt.data, eps_4.data); - grad_4[1].data = _mm512_fmadd_ps(grad_4[1].data, bias2_sqrt.data, eps_4.data); - grad_4[2].data = _mm512_fmadd_ps(grad_4[2].data, bias2_sqrt.data, eps_4.data); - grad_4[3].data = _mm512_fmadd_ps(grad_4[3].data, bias2_sqrt.data, eps_4.data); - grad_4[0].data = _mm512_div_ps(momntum_4[0].data, grad_4[0].data); - grad_4[1].data = _mm512_div_ps(momntum_4[1].data, grad_4[1].data); - grad_4[2].data = _mm512_div_ps(momntum_4[2].data, grad_4[2].data); - grad_4[3].data = _mm512_div_ps(momntum_4[3].data, grad_4[3].data); - - param_4[0].data = _mm512_fmadd_ps(grad_4[0].data, step_size_4.data, param_4[0].data); - param_4[1].data = _mm512_fmadd_ps(grad_4[1].data, step_size_4.data, param_4[1].data); - param_4[2].data = _mm512_fmadd_ps(grad_4[2].data, step_size_4.data, param_4[2].data); - param_4[3].data = _mm512_fmadd_ps(grad_4[3].data, step_size_4.data, param_4[3].data); - - _mm512_storeu_ps(_params + i, param_4[0].data); - _mm512_storeu_ps(_params + i + SIMD_WIDTH, param_4[1].data); - _mm512_storeu_ps(_params + i + (SIMD_WIDTH << 1), param_4[2].data); - _mm512_storeu_ps(_params + i + SIMD_WIDTH * 3, param_4[3].data); + momentum_4[0].data = SIMD_MUL(momentum_4[0].data, betta1_4.data); + momentum_4[0].data = + SIMD_FMA(grad_4[0].data, betta1_minus1_4.data, momentum_4[0].data); + momentum_4[1].data = SIMD_MUL(momentum_4[1].data, betta1_4.data); + momentum_4[1].data = + SIMD_FMA(grad_4[1].data, betta1_minus1_4.data, momentum_4[1].data); + momentum_4[2].data = SIMD_MUL(momentum_4[2].data, betta1_4.data); + momentum_4[2].data = + SIMD_FMA(grad_4[2].data, betta1_minus1_4.data, momentum_4[2].data); + momentum_4[3].data = SIMD_MUL(momentum_4[3].data, betta1_4.data); + momentum_4[3].data = + SIMD_FMA(grad_4[3].data, betta1_minus1_4.data, momentum_4[3].data); + + variance_4[0].data = SIMD_MUL(variance_4[0].data, betta2_4.data); + variance_4[1].data = SIMD_MUL(variance_4[1].data, betta2_4.data); + variance_4[2].data = SIMD_MUL(variance_4[2].data, betta2_4.data); + variance_4[3].data = SIMD_MUL(variance_4[3].data, betta2_4.data); + grad_4[0].data = SIMD_MUL(grad_4[0].data, grad_4[0].data); + grad_4[1].data = SIMD_MUL(grad_4[1].data, grad_4[1].data); + grad_4[2].data = SIMD_MUL(grad_4[2].data, grad_4[2].data); + grad_4[3].data = SIMD_MUL(grad_4[3].data, grad_4[3].data); + variance_4[0].data = + SIMD_FMA(grad_4[0].data, betta2_minus1_4.data, variance_4[0].data); + variance_4[1].data = + SIMD_FMA(grad_4[1].data, betta2_minus1_4.data, variance_4[1].data); + variance_4[2].data = + SIMD_FMA(grad_4[2].data, betta2_minus1_4.data, variance_4[2].data); + variance_4[3].data = + SIMD_FMA(grad_4[3].data, betta2_minus1_4.data, variance_4[3].data); + + grad_4[0].data = SIMD_SQRT(variance_4[0].data); + grad_4[1].data = SIMD_SQRT(variance_4[1].data); + grad_4[2].data = SIMD_SQRT(variance_4[2].data); + grad_4[3].data = SIMD_SQRT(variance_4[3].data); + + grad_4[0].data = SIMD_FMA(grad_4[0].data, bias2_sqrt.data, eps_4.data); + grad_4[1].data = SIMD_FMA(grad_4[1].data, bias2_sqrt.data, eps_4.data); + grad_4[2].data = SIMD_FMA(grad_4[2].data, bias2_sqrt.data, eps_4.data); + grad_4[3].data = SIMD_FMA(grad_4[3].data, bias2_sqrt.data, eps_4.data); + grad_4[0].data = SIMD_DIV(momentum_4[0].data, grad_4[0].data); + grad_4[1].data = SIMD_DIV(momentum_4[1].data, grad_4[1].data); + grad_4[2].data = SIMD_DIV(momentum_4[2].data, grad_4[2].data); + grad_4[3].data = SIMD_DIV(momentum_4[3].data, grad_4[3].data); + + param_4[0].data = SIMD_FMA(grad_4[0].data, step_size_4.data, param_4[0].data); + param_4[1].data = SIMD_FMA(grad_4[1].data, step_size_4.data, param_4[1].data); + param_4[2].data = SIMD_FMA(grad_4[2].data, step_size_4.data, param_4[2].data); + param_4[3].data = SIMD_FMA(grad_4[3].data, step_size_4.data, param_4[3].data); + + SIMD_STORE(_params + i, param_4[0].data); + SIMD_STORE(_params + i + SIMD_WIDTH, param_4[1].data); + SIMD_STORE(_params + i + (SIMD_WIDTH << 1), param_4[2].data); + SIMD_STORE(_params + i + SIMD_WIDTH * 3, param_4[3].data); if (dev_params) { - _mm512_storeu_ps(_doubled_buffer[_buf_index] + (i - t), param_4[0].data); - _mm512_storeu_ps(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH, + SIMD_STORE(_doubled_buffer[_buf_index] + (i - t), param_4[0].data); + SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH, param_4[1].data); - _mm512_storeu_ps(_doubled_buffer[_buf_index] + (i - t) + (SIMD_WIDTH << 1), + SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + (SIMD_WIDTH << 1), param_4[2].data); - _mm512_storeu_ps(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 3, + SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 3, param_4[3].data); } - _mm512_storeu_ps(_exp_avg + i, momntum_4[0].data); - _mm512_storeu_ps(_exp_avg + i + SIMD_WIDTH, momntum_4[1].data); - _mm512_storeu_ps(_exp_avg + i + (SIMD_WIDTH << 1), momntum_4[2].data); - _mm512_storeu_ps(_exp_avg + i + SIMD_WIDTH * 3, momntum_4[3].data); + SIMD_STORE(_exp_avg + i, momentum_4[0].data); + SIMD_STORE(_exp_avg + i + SIMD_WIDTH, momentum_4[1].data); + SIMD_STORE(_exp_avg + i + (SIMD_WIDTH << 1), momentum_4[2].data); + SIMD_STORE(_exp_avg + i + SIMD_WIDTH * 3, momentum_4[3].data); - _mm512_storeu_ps(_exp_avg_sq + i, varianc_4[0].data); - _mm512_storeu_ps(_exp_avg_sq + i + SIMD_WIDTH, varianc_4[1].data); - _mm512_storeu_ps(_exp_avg_sq + i + (SIMD_WIDTH << 1), varianc_4[2].data); - _mm512_storeu_ps(_exp_avg_sq + i + SIMD_WIDTH * 3, varianc_4[3].data); + SIMD_STORE(_exp_avg_sq + i, variance_4[0].data); + SIMD_STORE(_exp_avg_sq + i + SIMD_WIDTH, variance_4[1].data); + SIMD_STORE(_exp_avg_sq + i + (SIMD_WIDTH << 1), variance_4[2].data); + SIMD_STORE(_exp_avg_sq + i + SIMD_WIDTH * 3, variance_4[3].data); } - if (dev_params) { /* - #pragma omp parallel for - for (size_t j = 0; j < copy_size; j += 4) { - _doubled_buffer[_buf_index][j] = (__half)_params[t + j]; - _doubled_buffer[_buf_index][j + 1] = (__half)_params[t + j + 1]; - _doubled_buffer[_buf_index][j + 2] = (__half)_params[t + j + 2]; - _doubled_buffer[_buf_index][j + 3] = (__half)_params[t + j + 3]; - } - - CUDA_CHECK(cudaMemcpyAsync(dev_params + t, - _doubled_buffer[_buf_index], - copy_size * sizeof(__half), - cudaMemcpyHostToDevice, - Context::Instance().GetCurrentStream())); - */ + if (dev_params) { launch_param_update(_doubled_buffer[_buf_index], dev_params + t, copy_size, @@ -333,6 +313,7 @@ void Adam_Optimizer::Step_4(float* _params, _buf_index = !_buf_index; } } +#endif if (_param_size > rounded_size) Step((_params + rounded_size), (grads + rounded_size), @@ -352,9 +333,15 @@ int create_adam_optimizer(int optimizer_id, auto opt = std::make_shared(alpha, betta1, betta2, eps, weight_decay); s_optimizers[optimizer_id] = opt; - - std::cout << "Adam Optimizer #" << optimizer_id << " is created." << std::endl; - +#if defined(__AVX512__) + std::cout << "Adam Optimizer #" << optimizer_id << " is created with AVX512 arithmetic capability." << std::endl; +#else +#if defined(__AVX256__) + std::cout << "Adam Optimizer #" << optimizer_id << " is created with AVX2 arithmetic capability." << std::endl; +#else + std::cout << "Adam Optimizer #" << optimizer_id << " is created with scalar arithmetic capability." << std::endl; +#endif +#endif return 0; } @@ -365,35 +352,36 @@ void Adam_Optimizer::Step_8(float* _params, size_t _param_size, __half* dev_params) { - _betta1_t *= _betta1; - _betta2_t *= _betta2; + size_t rounded_size = 0; + +#if defined(__AVX512__) or defined(__AVX256__) - AVX_512 betta1_4; - betta1_4.data = _mm512_set1_ps(_betta1); - AVX_512 betta2_4; - betta2_4.data = _mm512_set1_ps(_betta2); + AVX_Data betta1_4; + betta1_4.data = SIMD_SET(_betta1); + AVX_Data betta2_4; + betta2_4.data = SIMD_SET(_betta2); float betta1_minus1 = 1 - _betta1; float betta2_minus1 = 1 - _betta2; - AVX_512 betta1_minus1_4; - betta1_minus1_4.data = _mm512_set1_ps(betta1_minus1); - AVX_512 betta2_minus1_4; - betta2_minus1_4.data = _mm512_set1_ps(betta2_minus1); + AVX_Data betta1_minus1_4; + betta1_minus1_4.data = SIMD_SET(betta1_minus1); + AVX_Data betta2_minus1_4; + betta2_minus1_4.data = SIMD_SET(betta2_minus1); float bias_correction1 = 1 - _betta1_t; float bias_correction2 = 1 / sqrt(1 - _betta2_t); - // AVX_512 bias_correction1_4 = _mm512_set1_ps(bias_correction1); - AVX_512 bias2_sqrt; - bias2_sqrt.data = _mm512_set1_ps(bias_correction2); + // AVX_Data bias_correction1_4 = SIMD_SET(bias_correction1); + AVX_Data bias2_sqrt; + bias2_sqrt.data = SIMD_SET(bias_correction2); - AVX_512 eps_4; - eps_4.data = _mm512_set1_ps(_eps); + AVX_Data eps_4; + eps_4.data = SIMD_SET(_eps); float step_size = -1 * _alpha / bias_correction1; - AVX_512 step_size_4; - step_size_4.data = _mm512_set1_ps(step_size); + AVX_Data step_size_4; + step_size_4.data = SIMD_SET(step_size); - size_t rounded_size = ROUND_DOWN(_param_size, (SIMD_WIDTH << 3)); + rounded_size = ROUND_DOWN(_param_size, (SIMD_WIDTH << 3)); for (size_t t = 0; t < rounded_size; t += TILE) { size_t copy_size = TILE; @@ -401,204 +389,204 @@ void Adam_Optimizer::Step_8(float* _params, size_t offset = copy_size + t; #pragma omp parallel for for (size_t i = t; i < offset; i += (SIMD_WIDTH << 3)) { - AVX_512 grad_4[8]; - grad_4[0].data = _mm512_loadu_ps(grads + i); - grad_4[1].data = _mm512_loadu_ps(grads + i + SIMD_WIDTH); - grad_4[2].data = _mm512_loadu_ps(grads + i + (SIMD_WIDTH << 1)); - grad_4[3].data = _mm512_loadu_ps(grads + i + SIMD_WIDTH * 3); - grad_4[4].data = _mm512_loadu_ps(grads + i + (SIMD_WIDTH << 2)); - grad_4[5].data = _mm512_loadu_ps(grads + i + SIMD_WIDTH * 5); - grad_4[6].data = _mm512_loadu_ps(grads + i + SIMD_WIDTH * 6); - grad_4[7].data = _mm512_loadu_ps(grads + i + SIMD_WIDTH * 7); - - AVX_512 momntum_4[8]; - momntum_4[0].data = _mm512_loadu_ps(_exp_avg + i); - momntum_4[1].data = _mm512_loadu_ps(_exp_avg + i + SIMD_WIDTH); - momntum_4[2].data = _mm512_loadu_ps(_exp_avg + i + (SIMD_WIDTH << 1)); - momntum_4[3].data = _mm512_loadu_ps(_exp_avg + i + SIMD_WIDTH * 3); - momntum_4[4].data = _mm512_loadu_ps(_exp_avg + i + (SIMD_WIDTH << 2)); - momntum_4[5].data = _mm512_loadu_ps(_exp_avg + i + SIMD_WIDTH * 5); - momntum_4[6].data = _mm512_loadu_ps(_exp_avg + i + SIMD_WIDTH * 6); - momntum_4[7].data = _mm512_loadu_ps(_exp_avg + i + SIMD_WIDTH * 7); - - AVX_512 varianc_4[8]; - varianc_4[0].data = _mm512_loadu_ps(_exp_avg_sq + i); - varianc_4[1].data = _mm512_loadu_ps(_exp_avg_sq + i + SIMD_WIDTH); - varianc_4[2].data = _mm512_loadu_ps(_exp_avg_sq + i + (SIMD_WIDTH << 1)); - varianc_4[3].data = _mm512_loadu_ps(_exp_avg_sq + i + SIMD_WIDTH * 3); - varianc_4[4].data = _mm512_loadu_ps(_exp_avg_sq + i + (SIMD_WIDTH << 2)); - varianc_4[5].data = _mm512_loadu_ps(_exp_avg_sq + i + SIMD_WIDTH * 5); - varianc_4[6].data = _mm512_loadu_ps(_exp_avg_sq + i + SIMD_WIDTH * 6); - varianc_4[7].data = _mm512_loadu_ps(_exp_avg_sq + i + SIMD_WIDTH * 7); - - AVX_512 param_4[8]; - param_4[0].data = _mm512_loadu_ps(_params + i); - param_4[1].data = _mm512_loadu_ps(_params + i + SIMD_WIDTH); - param_4[2].data = _mm512_loadu_ps(_params + i + (SIMD_WIDTH << 1)); - param_4[3].data = _mm512_loadu_ps(_params + i + SIMD_WIDTH * 3); - param_4[4].data = _mm512_loadu_ps(_params + i + (SIMD_WIDTH << 2)); - param_4[5].data = _mm512_loadu_ps(_params + i + SIMD_WIDTH * 5); - param_4[6].data = _mm512_loadu_ps(_params + i + SIMD_WIDTH * 6); - param_4[7].data = _mm512_loadu_ps(_params + i + SIMD_WIDTH * 7); + AVX_Data grad_4[8]; + grad_4[0].data = SIMD_LOAD(grads + i); + grad_4[1].data = SIMD_LOAD(grads + i + SIMD_WIDTH); + grad_4[2].data = SIMD_LOAD(grads + i + (SIMD_WIDTH << 1)); + grad_4[3].data = SIMD_LOAD(grads + i + SIMD_WIDTH * 3); + grad_4[4].data = SIMD_LOAD(grads + i + (SIMD_WIDTH << 2)); + grad_4[5].data = SIMD_LOAD(grads + i + SIMD_WIDTH * 5); + grad_4[6].data = SIMD_LOAD(grads + i + SIMD_WIDTH * 6); + grad_4[7].data = SIMD_LOAD(grads + i + SIMD_WIDTH * 7); + + AVX_Data momentum_4[8]; + momentum_4[0].data = SIMD_LOAD(_exp_avg + i); + momentum_4[1].data = SIMD_LOAD(_exp_avg + i + SIMD_WIDTH); + momentum_4[2].data = SIMD_LOAD(_exp_avg + i + (SIMD_WIDTH << 1)); + momentum_4[3].data = SIMD_LOAD(_exp_avg + i + SIMD_WIDTH * 3); + momentum_4[4].data = SIMD_LOAD(_exp_avg + i + (SIMD_WIDTH << 2)); + momentum_4[5].data = SIMD_LOAD(_exp_avg + i + SIMD_WIDTH * 5); + momentum_4[6].data = SIMD_LOAD(_exp_avg + i + SIMD_WIDTH * 6); + momentum_4[7].data = SIMD_LOAD(_exp_avg + i + SIMD_WIDTH * 7); + + AVX_Data variance_4[8]; + variance_4[0].data = SIMD_LOAD(_exp_avg_sq + i); + variance_4[1].data = SIMD_LOAD(_exp_avg_sq + i + SIMD_WIDTH); + variance_4[2].data = SIMD_LOAD(_exp_avg_sq + i + (SIMD_WIDTH << 1)); + variance_4[3].data = SIMD_LOAD(_exp_avg_sq + i + SIMD_WIDTH * 3); + variance_4[4].data = SIMD_LOAD(_exp_avg_sq + i + (SIMD_WIDTH << 2)); + variance_4[5].data = SIMD_LOAD(_exp_avg_sq + i + SIMD_WIDTH * 5); + variance_4[6].data = SIMD_LOAD(_exp_avg_sq + i + SIMD_WIDTH * 6); + variance_4[7].data = SIMD_LOAD(_exp_avg_sq + i + SIMD_WIDTH * 7); + + AVX_Data param_4[8]; + param_4[0].data = SIMD_LOAD(_params + i); + param_4[1].data = SIMD_LOAD(_params + i + SIMD_WIDTH); + param_4[2].data = SIMD_LOAD(_params + i + (SIMD_WIDTH << 1)); + param_4[3].data = SIMD_LOAD(_params + i + SIMD_WIDTH * 3); + param_4[4].data = SIMD_LOAD(_params + i + (SIMD_WIDTH << 2)); + param_4[5].data = SIMD_LOAD(_params + i + SIMD_WIDTH * 5); + param_4[6].data = SIMD_LOAD(_params + i + SIMD_WIDTH * 6); + param_4[7].data = SIMD_LOAD(_params + i + SIMD_WIDTH * 7); if (_weight_decay > 0) { - AVX_512 weight_decay4; - weight_decay4.data = _mm512_set1_ps(_weight_decay); + AVX_Data weight_decay4; + weight_decay4.data = SIMD_SET(_weight_decay); grad_4[0].data = - _mm512_fmadd_ps(param_4[0].data, weight_decay4.data, grad_4[0].data); + SIMD_FMA(param_4[0].data, weight_decay4.data, grad_4[0].data); grad_4[1].data = - _mm512_fmadd_ps(param_4[1].data, weight_decay4.data, grad_4[1].data); + SIMD_FMA(param_4[1].data, weight_decay4.data, grad_4[1].data); grad_4[2].data = - _mm512_fmadd_ps(param_4[2].data, weight_decay4.data, grad_4[2].data); + SIMD_FMA(param_4[2].data, weight_decay4.data, grad_4[2].data); grad_4[3].data = - _mm512_fmadd_ps(param_4[3].data, weight_decay4.data, grad_4[3].data); + SIMD_FMA(param_4[3].data, weight_decay4.data, grad_4[3].data); grad_4[4].data = - _mm512_fmadd_ps(param_4[4].data, weight_decay4.data, grad_4[4].data); + SIMD_FMA(param_4[4].data, weight_decay4.data, grad_4[4].data); grad_4[5].data = - _mm512_fmadd_ps(param_4[5].data, weight_decay4.data, grad_4[5].data); + SIMD_FMA(param_4[5].data, weight_decay4.data, grad_4[5].data); grad_4[6].data = - _mm512_fmadd_ps(param_4[6].data, weight_decay4.data, grad_4[6].data); + SIMD_FMA(param_4[6].data, weight_decay4.data, grad_4[6].data); grad_4[7].data = - _mm512_fmadd_ps(param_4[7].data, weight_decay4.data, grad_4[7].data); + SIMD_FMA(param_4[7].data, weight_decay4.data, grad_4[7].data); } - momntum_4[0].data = _mm512_mul_ps(momntum_4[0].data, betta1_4.data); - momntum_4[0].data = - _mm512_fmadd_ps(grad_4[0].data, betta1_minus1_4.data, momntum_4[0].data); - momntum_4[1].data = _mm512_mul_ps(momntum_4[1].data, betta1_4.data); - momntum_4[1].data = - _mm512_fmadd_ps(grad_4[1].data, betta1_minus1_4.data, momntum_4[1].data); - momntum_4[2].data = _mm512_mul_ps(momntum_4[2].data, betta1_4.data); - momntum_4[2].data = - _mm512_fmadd_ps(grad_4[2].data, betta1_minus1_4.data, momntum_4[2].data); - momntum_4[3].data = _mm512_mul_ps(momntum_4[3].data, betta1_4.data); - momntum_4[3].data = - _mm512_fmadd_ps(grad_4[3].data, betta1_minus1_4.data, momntum_4[3].data); - momntum_4[4].data = _mm512_mul_ps(momntum_4[4].data, betta1_4.data); - momntum_4[4].data = - _mm512_fmadd_ps(grad_4[4].data, betta1_minus1_4.data, momntum_4[4].data); - momntum_4[5].data = _mm512_mul_ps(momntum_4[5].data, betta1_4.data); - momntum_4[5].data = - _mm512_fmadd_ps(grad_4[5].data, betta1_minus1_4.data, momntum_4[5].data); - momntum_4[6].data = _mm512_mul_ps(momntum_4[6].data, betta1_4.data); - momntum_4[6].data = - _mm512_fmadd_ps(grad_4[6].data, betta1_minus1_4.data, momntum_4[6].data); - momntum_4[7].data = _mm512_mul_ps(momntum_4[7].data, betta1_4.data); - momntum_4[7].data = - _mm512_fmadd_ps(grad_4[7].data, betta1_minus1_4.data, momntum_4[7].data); - - varianc_4[0].data = _mm512_mul_ps(varianc_4[0].data, betta2_4.data); - varianc_4[1].data = _mm512_mul_ps(varianc_4[1].data, betta2_4.data); - varianc_4[2].data = _mm512_mul_ps(varianc_4[2].data, betta2_4.data); - varianc_4[3].data = _mm512_mul_ps(varianc_4[3].data, betta2_4.data); - varianc_4[4].data = _mm512_mul_ps(varianc_4[4].data, betta2_4.data); - varianc_4[5].data = _mm512_mul_ps(varianc_4[5].data, betta2_4.data); - varianc_4[6].data = _mm512_mul_ps(varianc_4[6].data, betta2_4.data); - varianc_4[7].data = _mm512_mul_ps(varianc_4[7].data, betta2_4.data); - grad_4[0].data = _mm512_mul_ps(grad_4[0].data, grad_4[0].data); - grad_4[1].data = _mm512_mul_ps(grad_4[1].data, grad_4[1].data); - grad_4[2].data = _mm512_mul_ps(grad_4[2].data, grad_4[2].data); - grad_4[3].data = _mm512_mul_ps(grad_4[3].data, grad_4[3].data); - grad_4[4].data = _mm512_mul_ps(grad_4[4].data, grad_4[4].data); - grad_4[5].data = _mm512_mul_ps(grad_4[5].data, grad_4[5].data); - grad_4[6].data = _mm512_mul_ps(grad_4[6].data, grad_4[6].data); - grad_4[7].data = _mm512_mul_ps(grad_4[7].data, grad_4[7].data); - varianc_4[0].data = - _mm512_fmadd_ps(grad_4[0].data, betta2_minus1_4.data, varianc_4[0].data); - varianc_4[1].data = - _mm512_fmadd_ps(grad_4[1].data, betta2_minus1_4.data, varianc_4[1].data); - varianc_4[2].data = - _mm512_fmadd_ps(grad_4[2].data, betta2_minus1_4.data, varianc_4[2].data); - varianc_4[3].data = - _mm512_fmadd_ps(grad_4[3].data, betta2_minus1_4.data, varianc_4[3].data); - varianc_4[4].data = - _mm512_fmadd_ps(grad_4[4].data, betta2_minus1_4.data, varianc_4[4].data); - varianc_4[5].data = - _mm512_fmadd_ps(grad_4[5].data, betta2_minus1_4.data, varianc_4[5].data); - varianc_4[6].data = - _mm512_fmadd_ps(grad_4[6].data, betta2_minus1_4.data, varianc_4[6].data); - varianc_4[7].data = - _mm512_fmadd_ps(grad_4[7].data, betta2_minus1_4.data, varianc_4[7].data); - - grad_4[0].data = _mm512_sqrt_ps(varianc_4[0].data); - grad_4[1].data = _mm512_sqrt_ps(varianc_4[1].data); - grad_4[2].data = _mm512_sqrt_ps(varianc_4[2].data); - grad_4[3].data = _mm512_sqrt_ps(varianc_4[3].data); - grad_4[4].data = _mm512_sqrt_ps(varianc_4[4].data); - grad_4[5].data = _mm512_sqrt_ps(varianc_4[5].data); - grad_4[6].data = _mm512_sqrt_ps(varianc_4[6].data); - grad_4[7].data = _mm512_sqrt_ps(varianc_4[7].data); - - grad_4[0].data = _mm512_fmadd_ps(grad_4[0].data, bias2_sqrt.data, eps_4.data); - grad_4[1].data = _mm512_fmadd_ps(grad_4[1].data, bias2_sqrt.data, eps_4.data); - grad_4[2].data = _mm512_fmadd_ps(grad_4[2].data, bias2_sqrt.data, eps_4.data); - grad_4[3].data = _mm512_fmadd_ps(grad_4[3].data, bias2_sqrt.data, eps_4.data); - grad_4[4].data = _mm512_fmadd_ps(grad_4[4].data, bias2_sqrt.data, eps_4.data); - grad_4[5].data = _mm512_fmadd_ps(grad_4[5].data, bias2_sqrt.data, eps_4.data); - grad_4[6].data = _mm512_fmadd_ps(grad_4[6].data, bias2_sqrt.data, eps_4.data); - grad_4[7].data = _mm512_fmadd_ps(grad_4[7].data, bias2_sqrt.data, eps_4.data); - grad_4[0].data = _mm512_div_ps(momntum_4[0].data, grad_4[0].data); - grad_4[1].data = _mm512_div_ps(momntum_4[1].data, grad_4[1].data); - grad_4[2].data = _mm512_div_ps(momntum_4[2].data, grad_4[2].data); - grad_4[3].data = _mm512_div_ps(momntum_4[3].data, grad_4[3].data); - grad_4[4].data = _mm512_div_ps(momntum_4[4].data, grad_4[4].data); - grad_4[5].data = _mm512_div_ps(momntum_4[5].data, grad_4[5].data); - grad_4[6].data = _mm512_div_ps(momntum_4[6].data, grad_4[6].data); - grad_4[7].data = _mm512_div_ps(momntum_4[7].data, grad_4[7].data); - - param_4[0].data = _mm512_fmadd_ps(grad_4[0].data, step_size_4.data, param_4[0].data); - param_4[1].data = _mm512_fmadd_ps(grad_4[1].data, step_size_4.data, param_4[1].data); - param_4[2].data = _mm512_fmadd_ps(grad_4[2].data, step_size_4.data, param_4[2].data); - param_4[3].data = _mm512_fmadd_ps(grad_4[3].data, step_size_4.data, param_4[3].data); - param_4[4].data = _mm512_fmadd_ps(grad_4[4].data, step_size_4.data, param_4[4].data); - param_4[5].data = _mm512_fmadd_ps(grad_4[5].data, step_size_4.data, param_4[5].data); - param_4[6].data = _mm512_fmadd_ps(grad_4[6].data, step_size_4.data, param_4[6].data); - param_4[7].data = _mm512_fmadd_ps(grad_4[7].data, step_size_4.data, param_4[7].data); - - _mm512_storeu_ps(_params + i, param_4[0].data); - _mm512_storeu_ps(_params + i + SIMD_WIDTH, param_4[1].data); - _mm512_storeu_ps(_params + i + (SIMD_WIDTH << 1), param_4[2].data); - _mm512_storeu_ps(_params + i + SIMD_WIDTH * 3, param_4[3].data); - _mm512_storeu_ps(_params + i + (SIMD_WIDTH << 2), param_4[4].data); - _mm512_storeu_ps(_params + i + SIMD_WIDTH * 5, param_4[5].data); - _mm512_storeu_ps(_params + i + SIMD_WIDTH * 6, param_4[6].data); - _mm512_storeu_ps(_params + i + SIMD_WIDTH * 7, param_4[7].data); + momentum_4[0].data = SIMD_MUL(momentum_4[0].data, betta1_4.data); + momentum_4[0].data = + SIMD_FMA(grad_4[0].data, betta1_minus1_4.data, momentum_4[0].data); + momentum_4[1].data = SIMD_MUL(momentum_4[1].data, betta1_4.data); + momentum_4[1].data = + SIMD_FMA(grad_4[1].data, betta1_minus1_4.data, momentum_4[1].data); + momentum_4[2].data = SIMD_MUL(momentum_4[2].data, betta1_4.data); + momentum_4[2].data = + SIMD_FMA(grad_4[2].data, betta1_minus1_4.data, momentum_4[2].data); + momentum_4[3].data = SIMD_MUL(momentum_4[3].data, betta1_4.data); + momentum_4[3].data = + SIMD_FMA(grad_4[3].data, betta1_minus1_4.data, momentum_4[3].data); + momentum_4[4].data = SIMD_MUL(momentum_4[4].data, betta1_4.data); + momentum_4[4].data = + SIMD_FMA(grad_4[4].data, betta1_minus1_4.data, momentum_4[4].data); + momentum_4[5].data = SIMD_MUL(momentum_4[5].data, betta1_4.data); + momentum_4[5].data = + SIMD_FMA(grad_4[5].data, betta1_minus1_4.data, momentum_4[5].data); + momentum_4[6].data = SIMD_MUL(momentum_4[6].data, betta1_4.data); + momentum_4[6].data = + SIMD_FMA(grad_4[6].data, betta1_minus1_4.data, momentum_4[6].data); + momentum_4[7].data = SIMD_MUL(momentum_4[7].data, betta1_4.data); + momentum_4[7].data = + SIMD_FMA(grad_4[7].data, betta1_minus1_4.data, momentum_4[7].data); + + variance_4[0].data = SIMD_MUL(variance_4[0].data, betta2_4.data); + variance_4[1].data = SIMD_MUL(variance_4[1].data, betta2_4.data); + variance_4[2].data = SIMD_MUL(variance_4[2].data, betta2_4.data); + variance_4[3].data = SIMD_MUL(variance_4[3].data, betta2_4.data); + variance_4[4].data = SIMD_MUL(variance_4[4].data, betta2_4.data); + variance_4[5].data = SIMD_MUL(variance_4[5].data, betta2_4.data); + variance_4[6].data = SIMD_MUL(variance_4[6].data, betta2_4.data); + variance_4[7].data = SIMD_MUL(variance_4[7].data, betta2_4.data); + grad_4[0].data = SIMD_MUL(grad_4[0].data, grad_4[0].data); + grad_4[1].data = SIMD_MUL(grad_4[1].data, grad_4[1].data); + grad_4[2].data = SIMD_MUL(grad_4[2].data, grad_4[2].data); + grad_4[3].data = SIMD_MUL(grad_4[3].data, grad_4[3].data); + grad_4[4].data = SIMD_MUL(grad_4[4].data, grad_4[4].data); + grad_4[5].data = SIMD_MUL(grad_4[5].data, grad_4[5].data); + grad_4[6].data = SIMD_MUL(grad_4[6].data, grad_4[6].data); + grad_4[7].data = SIMD_MUL(grad_4[7].data, grad_4[7].data); + variance_4[0].data = + SIMD_FMA(grad_4[0].data, betta2_minus1_4.data, variance_4[0].data); + variance_4[1].data = + SIMD_FMA(grad_4[1].data, betta2_minus1_4.data, variance_4[1].data); + variance_4[2].data = + SIMD_FMA(grad_4[2].data, betta2_minus1_4.data, variance_4[2].data); + variance_4[3].data = + SIMD_FMA(grad_4[3].data, betta2_minus1_4.data, variance_4[3].data); + variance_4[4].data = + SIMD_FMA(grad_4[4].data, betta2_minus1_4.data, variance_4[4].data); + variance_4[5].data = + SIMD_FMA(grad_4[5].data, betta2_minus1_4.data, variance_4[5].data); + variance_4[6].data = + SIMD_FMA(grad_4[6].data, betta2_minus1_4.data, variance_4[6].data); + variance_4[7].data = + SIMD_FMA(grad_4[7].data, betta2_minus1_4.data, variance_4[7].data); + + grad_4[0].data = SIMD_SQRT(variance_4[0].data); + grad_4[1].data = SIMD_SQRT(variance_4[1].data); + grad_4[2].data = SIMD_SQRT(variance_4[2].data); + grad_4[3].data = SIMD_SQRT(variance_4[3].data); + grad_4[4].data = SIMD_SQRT(variance_4[4].data); + grad_4[5].data = SIMD_SQRT(variance_4[5].data); + grad_4[6].data = SIMD_SQRT(variance_4[6].data); + grad_4[7].data = SIMD_SQRT(variance_4[7].data); + + grad_4[0].data = SIMD_FMA(grad_4[0].data, bias2_sqrt.data, eps_4.data); + grad_4[1].data = SIMD_FMA(grad_4[1].data, bias2_sqrt.data, eps_4.data); + grad_4[2].data = SIMD_FMA(grad_4[2].data, bias2_sqrt.data, eps_4.data); + grad_4[3].data = SIMD_FMA(grad_4[3].data, bias2_sqrt.data, eps_4.data); + grad_4[4].data = SIMD_FMA(grad_4[4].data, bias2_sqrt.data, eps_4.data); + grad_4[5].data = SIMD_FMA(grad_4[5].data, bias2_sqrt.data, eps_4.data); + grad_4[6].data = SIMD_FMA(grad_4[6].data, bias2_sqrt.data, eps_4.data); + grad_4[7].data = SIMD_FMA(grad_4[7].data, bias2_sqrt.data, eps_4.data); + grad_4[0].data = SIMD_DIV(momentum_4[0].data, grad_4[0].data); + grad_4[1].data = SIMD_DIV(momentum_4[1].data, grad_4[1].data); + grad_4[2].data = SIMD_DIV(momentum_4[2].data, grad_4[2].data); + grad_4[3].data = SIMD_DIV(momentum_4[3].data, grad_4[3].data); + grad_4[4].data = SIMD_DIV(momentum_4[4].data, grad_4[4].data); + grad_4[5].data = SIMD_DIV(momentum_4[5].data, grad_4[5].data); + grad_4[6].data = SIMD_DIV(momentum_4[6].data, grad_4[6].data); + grad_4[7].data = SIMD_DIV(momentum_4[7].data, grad_4[7].data); + + param_4[0].data = SIMD_FMA(grad_4[0].data, step_size_4.data, param_4[0].data); + param_4[1].data = SIMD_FMA(grad_4[1].data, step_size_4.data, param_4[1].data); + param_4[2].data = SIMD_FMA(grad_4[2].data, step_size_4.data, param_4[2].data); + param_4[3].data = SIMD_FMA(grad_4[3].data, step_size_4.data, param_4[3].data); + param_4[4].data = SIMD_FMA(grad_4[4].data, step_size_4.data, param_4[4].data); + param_4[5].data = SIMD_FMA(grad_4[5].data, step_size_4.data, param_4[5].data); + param_4[6].data = SIMD_FMA(grad_4[6].data, step_size_4.data, param_4[6].data); + param_4[7].data = SIMD_FMA(grad_4[7].data, step_size_4.data, param_4[7].data); + + SIMD_STORE(_params + i, param_4[0].data); + SIMD_STORE(_params + i + SIMD_WIDTH, param_4[1].data); + SIMD_STORE(_params + i + (SIMD_WIDTH << 1), param_4[2].data); + SIMD_STORE(_params + i + SIMD_WIDTH * 3, param_4[3].data); + SIMD_STORE(_params + i + (SIMD_WIDTH << 2), param_4[4].data); + SIMD_STORE(_params + i + SIMD_WIDTH * 5, param_4[5].data); + SIMD_STORE(_params + i + SIMD_WIDTH * 6, param_4[6].data); + SIMD_STORE(_params + i + SIMD_WIDTH * 7, param_4[7].data); if (dev_params) { - _mm512_storeu_ps(_doubled_buffer[_buf_index] + (i - t), param_4[0].data); - _mm512_storeu_ps(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH, + SIMD_STORE(_doubled_buffer[_buf_index] + (i - t), param_4[0].data); + SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH, param_4[1].data); - _mm512_storeu_ps(_doubled_buffer[_buf_index] + (i - t) + (SIMD_WIDTH << 1), + SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + (SIMD_WIDTH << 1), param_4[2].data); - _mm512_storeu_ps(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 3, + SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 3, param_4[3].data); - _mm512_storeu_ps(_doubled_buffer[_buf_index] + (i - t) + (SIMD_WIDTH << 2), + SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + (SIMD_WIDTH << 2), param_4[4].data); - _mm512_storeu_ps(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 5, + SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 5, param_4[5].data); - _mm512_storeu_ps(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 6, + SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 6, param_4[6].data); - _mm512_storeu_ps(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 7, + SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 7, param_4[7].data); } - _mm512_storeu_ps(_exp_avg + i, momntum_4[0].data); - _mm512_storeu_ps(_exp_avg + i + SIMD_WIDTH, momntum_4[1].data); - _mm512_storeu_ps(_exp_avg + i + (SIMD_WIDTH << 1), momntum_4[2].data); - _mm512_storeu_ps(_exp_avg + i + SIMD_WIDTH * 3, momntum_4[3].data); - _mm512_storeu_ps(_exp_avg + i + (SIMD_WIDTH << 2), momntum_4[4].data); - _mm512_storeu_ps(_exp_avg + i + SIMD_WIDTH * 5, momntum_4[5].data); - _mm512_storeu_ps(_exp_avg + i + SIMD_WIDTH * 6, momntum_4[6].data); - _mm512_storeu_ps(_exp_avg + i + SIMD_WIDTH * 7, momntum_4[7].data); - - _mm512_storeu_ps(_exp_avg_sq + i, varianc_4[0].data); - _mm512_storeu_ps(_exp_avg_sq + i + SIMD_WIDTH, varianc_4[1].data); - _mm512_storeu_ps(_exp_avg_sq + i + (SIMD_WIDTH << 1), varianc_4[2].data); - _mm512_storeu_ps(_exp_avg_sq + i + SIMD_WIDTH * 3, varianc_4[3].data); - _mm512_storeu_ps(_exp_avg_sq + i + (SIMD_WIDTH << 2), varianc_4[4].data); - _mm512_storeu_ps(_exp_avg_sq + i + SIMD_WIDTH * 5, varianc_4[5].data); - _mm512_storeu_ps(_exp_avg_sq + i + SIMD_WIDTH * 6, varianc_4[6].data); - _mm512_storeu_ps(_exp_avg_sq + i + SIMD_WIDTH * 7, varianc_4[7].data); + SIMD_STORE(_exp_avg + i, momentum_4[0].data); + SIMD_STORE(_exp_avg + i + SIMD_WIDTH, momentum_4[1].data); + SIMD_STORE(_exp_avg + i + (SIMD_WIDTH << 1), momentum_4[2].data); + SIMD_STORE(_exp_avg + i + SIMD_WIDTH * 3, momentum_4[3].data); + SIMD_STORE(_exp_avg + i + (SIMD_WIDTH << 2), momentum_4[4].data); + SIMD_STORE(_exp_avg + i + SIMD_WIDTH * 5, momentum_4[5].data); + SIMD_STORE(_exp_avg + i + SIMD_WIDTH * 6, momentum_4[6].data); + SIMD_STORE(_exp_avg + i + SIMD_WIDTH * 7, momentum_4[7].data); + + SIMD_STORE(_exp_avg_sq + i, variance_4[0].data); + SIMD_STORE(_exp_avg_sq + i + SIMD_WIDTH, variance_4[1].data); + SIMD_STORE(_exp_avg_sq + i + (SIMD_WIDTH << 1), variance_4[2].data); + SIMD_STORE(_exp_avg_sq + i + SIMD_WIDTH * 3, variance_4[3].data); + SIMD_STORE(_exp_avg_sq + i + (SIMD_WIDTH << 2), variance_4[4].data); + SIMD_STORE(_exp_avg_sq + i + SIMD_WIDTH * 5, variance_4[5].data); + SIMD_STORE(_exp_avg_sq + i + SIMD_WIDTH * 6, variance_4[6].data); + SIMD_STORE(_exp_avg_sq + i + SIMD_WIDTH * 7, variance_4[7].data); } if (dev_params) { launch_param_update(_doubled_buffer[_buf_index], @@ -608,6 +596,7 @@ void Adam_Optimizer::Step_8(float* _params, _buf_index = !_buf_index; } } +#endif if (_param_size > rounded_size) Step_4((_params + rounded_size), (grads + rounded_size), @@ -635,7 +624,7 @@ int ds_adam_step(int optimizer_id, std::shared_ptr opt = std::static_pointer_cast(s_optimizers[optimizer_id]); - + opt->IncrementStep(); opt->Step_8(params_ptr, grads_ptr, exp_avg_ptr, exp_avg_sq_ptr, params_c.size(0)); return 0; @@ -662,7 +651,7 @@ int ds_adam_step_plus_copy(int optimizer_id, std::shared_ptr opt = std::static_pointer_cast(s_optimizers[optimizer_id]); - + opt->IncrementStep(); opt->Step_8( params_ptr, grads_ptr, exp_avg_ptr, exp_avg_sq_ptr, params_c.size(0), gpu_params_ptr); diff --git a/csrc/includes/cpu_adam.h b/csrc/includes/cpu_adam.h old mode 100644 new mode 100755 index b373162612f2..da1f73f3e44f --- a/csrc/includes/cpu_adam.h +++ b/csrc/includes/cpu_adam.h @@ -8,6 +8,8 @@ #include "cublas_v2.h" #include "cuda.h" #include "curand.h" +#include +#include #define CUDA_CHECK(callstr) \ { \ @@ -18,10 +20,30 @@ } \ } -#define SIMD_WIDTH 16 - #define TILE (1024 * 1024 * 1024) +#if defined(__AVX512__) + #define SIMD_STORE(a, d) _mm512_storeu_ps(a, d) + #define SIMD_LOAD(x) _mm512_loadu_ps(x) + #define SIMD_SET(x) _mm512_set1_ps(x) + #define SIMD_MUL(x, y) _mm512_mul_ps(x, y) + #define SIMD_FMA(x, y, c) _mm512_fmadd_ps(x, y, c) + #define SIMD_SQRT(x) _mm512_sqrt_ps(x) + #define SIMD_DIV(x, y) _mm512_div_ps(x, y) + #define SIMD_WIDTH 16 +#else + #if defined(__AVX256__) + #define SIMD_STORE(a, d) _mm256_storeu_ps(a, d) + #define SIMD_LOAD(x) _mm256_loadu_ps(x) + #define SIMD_SET(x) _mm256_set1_ps(x) + #define SIMD_MUL(x, y) _mm256_mul_ps(x, y) + #define SIMD_FMA(x, y, c) _mm256_fmadd_ps(x, y, c) + #define SIMD_SQRT(x) _mm256_sqrt_ps(x) + #define SIMD_DIV(x, y) _mm256_div_ps(x, y) + #define SIMD_WIDTH 8 + #endif +#endif + class Adam_Optimizer { public: Adam_Optimizer(float alpha = 1e-3, @@ -64,12 +86,22 @@ class Adam_Optimizer { float* _exp_avg_sq, size_t _param_size, __half* dev_params = nullptr); - + inline void IncrementStep() + { + _betta1_t *= _betta1; + _betta2_t *= _betta2; + } private: - union AVX_512 { +#if defined(__AVX512__) or defined(__AVX256__) + union AVX_Data { +#if defined(__AVX512__) __m512 data; +#else + __m256 data; +#endif // float data_f[16]; }; +#endif float _alpha; float _betta1; diff --git a/setup.py b/setup.py index eb8d6ed07e51..f59e71adc2c7 100755 --- a/setup.py +++ b/setup.py @@ -96,6 +96,17 @@ def fetch_requirements(path): version_ge_1_5 = ['-DVERSION_GE_1_5'] version_dependent_macros = version_ge_1_1 + version_ge_1_3 + version_ge_1_5 +import cpufeature +d = cpufeature.CPUFeature + +SIMD_WIDTH = '' +if d['AVX512f']: + SIMD_WIDTH = '-D__AVX512__' +elif d['AVX2']: + SIMD_WIDTH = '-D__AVX256__' +print("SIMD_WIDTH = ", SIMD_WIDTH) + + ext_modules = [] ## Lamb ## @@ -135,7 +146,8 @@ def fetch_requirements(path): '-g', '-Wno-reorder', '-march=native', - '-fopenmp' + '-fopenmp', + SIMD_WIDTH ], 'nvcc': [ '-O3', From c4d304203950bfef95c6541a48a7bf4c1aff4fbf Mon Sep 17 00:00:00 2001 From: Reza Yazdani Date: Tue, 8 Sep 2020 21:40:10 +0000 Subject: [PATCH 2/8] running precommit --- csrc/adam/cpu_adam.cpp | 159 +++++++++++++++------------------------ csrc/includes/cpu_adam.h | 47 ++++++------ setup.py | 1 - 3 files changed, 83 insertions(+), 124 deletions(-) mode change 100755 => 100644 csrc/adam/cpu_adam.cpp mode change 100755 => 100644 csrc/includes/cpu_adam.h diff --git a/csrc/adam/cpu_adam.cpp b/csrc/adam/cpu_adam.cpp old mode 100755 new mode 100644 index bc0c32e45f52..380bc4ea0ab0 --- a/csrc/adam/cpu_adam.cpp +++ b/csrc/adam/cpu_adam.cpp @@ -46,7 +46,7 @@ void Adam_Optimizer::Step(float* _params, betta1_minus1_4.data = SIMD_SET(betta1_minus1); AVX_Data betta2_minus1_4; betta2_minus1_4.data = SIMD_SET(betta2_minus1); - + AVX_Data bias2_sqrt; bias2_sqrt.data = SIMD_SET(bias_correction2); @@ -57,8 +57,7 @@ void Adam_Optimizer::Step(float* _params, step_size_4.data = SIMD_SET(step_size); AVX_Data weight_decay4; - if (_weight_decay > 0) - weight_decay4.data = SIMD_SET(_weight_decay); + if (_weight_decay > 0) weight_decay4.data = SIMD_SET(_weight_decay); rounded_size = ROUND_DOWN(_param_size, SIMD_WIDTH); @@ -120,7 +119,7 @@ void Adam_Optimizer::Step(float* _params, float param = _params[k]; float momentum = _exp_avg[k]; float variance = _exp_avg_sq[k]; - if (_weight_decay > 0) grad = param * _weight_decay + grad; + if (_weight_decay > 0) grad = param * _weight_decay + grad; momentum *= momentum * _betta1; momentum = grad * betta1_minus1 + momentum; @@ -157,8 +156,8 @@ void Adam_Optimizer::Step_4(float* _params, __half* dev_params) { size_t rounded_size = 0; - -#if defined(__AVX512__) or defined(__AVX256__) + +#if defined(__AVX512__) or defined(__AVX256__) AVX_Data betta1_4; betta1_4.data = SIMD_SET(_betta1); @@ -220,28 +219,20 @@ void Adam_Optimizer::Step_4(float* _params, if (_weight_decay > 0) { AVX_Data weight_decay4; weight_decay4.data = SIMD_SET(_weight_decay); - grad_4[0].data = - SIMD_FMA(param_4[0].data, weight_decay4.data, grad_4[0].data); - grad_4[1].data = - SIMD_FMA(param_4[1].data, weight_decay4.data, grad_4[1].data); - grad_4[2].data = - SIMD_FMA(param_4[2].data, weight_decay4.data, grad_4[2].data); - grad_4[3].data = - SIMD_FMA(param_4[3].data, weight_decay4.data, grad_4[3].data); + grad_4[0].data = SIMD_FMA(param_4[0].data, weight_decay4.data, grad_4[0].data); + grad_4[1].data = SIMD_FMA(param_4[1].data, weight_decay4.data, grad_4[1].data); + grad_4[2].data = SIMD_FMA(param_4[2].data, weight_decay4.data, grad_4[2].data); + grad_4[3].data = SIMD_FMA(param_4[3].data, weight_decay4.data, grad_4[3].data); } momentum_4[0].data = SIMD_MUL(momentum_4[0].data, betta1_4.data); - momentum_4[0].data = - SIMD_FMA(grad_4[0].data, betta1_minus1_4.data, momentum_4[0].data); + momentum_4[0].data = SIMD_FMA(grad_4[0].data, betta1_minus1_4.data, momentum_4[0].data); momentum_4[1].data = SIMD_MUL(momentum_4[1].data, betta1_4.data); - momentum_4[1].data = - SIMD_FMA(grad_4[1].data, betta1_minus1_4.data, momentum_4[1].data); + momentum_4[1].data = SIMD_FMA(grad_4[1].data, betta1_minus1_4.data, momentum_4[1].data); momentum_4[2].data = SIMD_MUL(momentum_4[2].data, betta1_4.data); - momentum_4[2].data = - SIMD_FMA(grad_4[2].data, betta1_minus1_4.data, momentum_4[2].data); + momentum_4[2].data = SIMD_FMA(grad_4[2].data, betta1_minus1_4.data, momentum_4[2].data); momentum_4[3].data = SIMD_MUL(momentum_4[3].data, betta1_4.data); - momentum_4[3].data = - SIMD_FMA(grad_4[3].data, betta1_minus1_4.data, momentum_4[3].data); + momentum_4[3].data = SIMD_FMA(grad_4[3].data, betta1_minus1_4.data, momentum_4[3].data); variance_4[0].data = SIMD_MUL(variance_4[0].data, betta2_4.data); variance_4[1].data = SIMD_MUL(variance_4[1].data, betta2_4.data); @@ -251,14 +242,10 @@ void Adam_Optimizer::Step_4(float* _params, grad_4[1].data = SIMD_MUL(grad_4[1].data, grad_4[1].data); grad_4[2].data = SIMD_MUL(grad_4[2].data, grad_4[2].data); grad_4[3].data = SIMD_MUL(grad_4[3].data, grad_4[3].data); - variance_4[0].data = - SIMD_FMA(grad_4[0].data, betta2_minus1_4.data, variance_4[0].data); - variance_4[1].data = - SIMD_FMA(grad_4[1].data, betta2_minus1_4.data, variance_4[1].data); - variance_4[2].data = - SIMD_FMA(grad_4[2].data, betta2_minus1_4.data, variance_4[2].data); - variance_4[3].data = - SIMD_FMA(grad_4[3].data, betta2_minus1_4.data, variance_4[3].data); + variance_4[0].data = SIMD_FMA(grad_4[0].data, betta2_minus1_4.data, variance_4[0].data); + variance_4[1].data = SIMD_FMA(grad_4[1].data, betta2_minus1_4.data, variance_4[1].data); + variance_4[2].data = SIMD_FMA(grad_4[2].data, betta2_minus1_4.data, variance_4[2].data); + variance_4[3].data = SIMD_FMA(grad_4[3].data, betta2_minus1_4.data, variance_4[3].data); grad_4[0].data = SIMD_SQRT(variance_4[0].data); grad_4[1].data = SIMD_SQRT(variance_4[1].data); @@ -286,12 +273,10 @@ void Adam_Optimizer::Step_4(float* _params, if (dev_params) { SIMD_STORE(_doubled_buffer[_buf_index] + (i - t), param_4[0].data); - SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH, - param_4[1].data); + SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH, param_4[1].data); SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + (SIMD_WIDTH << 1), - param_4[2].data); - SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 3, - param_4[3].data); + param_4[2].data); + SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 3, param_4[3].data); } SIMD_STORE(_exp_avg + i, momentum_4[0].data); @@ -334,12 +319,15 @@ int create_adam_optimizer(int optimizer_id, s_optimizers[optimizer_id] = opt; #if defined(__AVX512__) - std::cout << "Adam Optimizer #" << optimizer_id << " is created with AVX512 arithmetic capability." << std::endl; + std::cout << "Adam Optimizer #" << optimizer_id + << " is created with AVX512 arithmetic capability." << std::endl; #else #if defined(__AVX256__) - std::cout << "Adam Optimizer #" << optimizer_id << " is created with AVX2 arithmetic capability." << std::endl; + std::cout << "Adam Optimizer #" << optimizer_id + << " is created with AVX2 arithmetic capability." << std::endl; #else - std::cout << "Adam Optimizer #" << optimizer_id << " is created with scalar arithmetic capability." << std::endl; + std::cout << "Adam Optimizer #" << optimizer_id + << " is created with scalar arithmetic capability." << std::endl; #endif #endif return 0; @@ -353,8 +341,8 @@ void Adam_Optimizer::Step_8(float* _params, __half* dev_params) { size_t rounded_size = 0; - -#if defined(__AVX512__) or defined(__AVX256__) + +#if defined(__AVX512__) or defined(__AVX256__) AVX_Data betta1_4; betta1_4.data = SIMD_SET(_betta1); @@ -432,48 +420,32 @@ void Adam_Optimizer::Step_8(float* _params, if (_weight_decay > 0) { AVX_Data weight_decay4; weight_decay4.data = SIMD_SET(_weight_decay); - grad_4[0].data = - SIMD_FMA(param_4[0].data, weight_decay4.data, grad_4[0].data); - grad_4[1].data = - SIMD_FMA(param_4[1].data, weight_decay4.data, grad_4[1].data); - grad_4[2].data = - SIMD_FMA(param_4[2].data, weight_decay4.data, grad_4[2].data); - grad_4[3].data = - SIMD_FMA(param_4[3].data, weight_decay4.data, grad_4[3].data); - grad_4[4].data = - SIMD_FMA(param_4[4].data, weight_decay4.data, grad_4[4].data); - grad_4[5].data = - SIMD_FMA(param_4[5].data, weight_decay4.data, grad_4[5].data); - grad_4[6].data = - SIMD_FMA(param_4[6].data, weight_decay4.data, grad_4[6].data); - grad_4[7].data = - SIMD_FMA(param_4[7].data, weight_decay4.data, grad_4[7].data); + grad_4[0].data = SIMD_FMA(param_4[0].data, weight_decay4.data, grad_4[0].data); + grad_4[1].data = SIMD_FMA(param_4[1].data, weight_decay4.data, grad_4[1].data); + grad_4[2].data = SIMD_FMA(param_4[2].data, weight_decay4.data, grad_4[2].data); + grad_4[3].data = SIMD_FMA(param_4[3].data, weight_decay4.data, grad_4[3].data); + grad_4[4].data = SIMD_FMA(param_4[4].data, weight_decay4.data, grad_4[4].data); + grad_4[5].data = SIMD_FMA(param_4[5].data, weight_decay4.data, grad_4[5].data); + grad_4[6].data = SIMD_FMA(param_4[6].data, weight_decay4.data, grad_4[6].data); + grad_4[7].data = SIMD_FMA(param_4[7].data, weight_decay4.data, grad_4[7].data); } momentum_4[0].data = SIMD_MUL(momentum_4[0].data, betta1_4.data); - momentum_4[0].data = - SIMD_FMA(grad_4[0].data, betta1_minus1_4.data, momentum_4[0].data); + momentum_4[0].data = SIMD_FMA(grad_4[0].data, betta1_minus1_4.data, momentum_4[0].data); momentum_4[1].data = SIMD_MUL(momentum_4[1].data, betta1_4.data); - momentum_4[1].data = - SIMD_FMA(grad_4[1].data, betta1_minus1_4.data, momentum_4[1].data); + momentum_4[1].data = SIMD_FMA(grad_4[1].data, betta1_minus1_4.data, momentum_4[1].data); momentum_4[2].data = SIMD_MUL(momentum_4[2].data, betta1_4.data); - momentum_4[2].data = - SIMD_FMA(grad_4[2].data, betta1_minus1_4.data, momentum_4[2].data); + momentum_4[2].data = SIMD_FMA(grad_4[2].data, betta1_minus1_4.data, momentum_4[2].data); momentum_4[3].data = SIMD_MUL(momentum_4[3].data, betta1_4.data); - momentum_4[3].data = - SIMD_FMA(grad_4[3].data, betta1_minus1_4.data, momentum_4[3].data); + momentum_4[3].data = SIMD_FMA(grad_4[3].data, betta1_minus1_4.data, momentum_4[3].data); momentum_4[4].data = SIMD_MUL(momentum_4[4].data, betta1_4.data); - momentum_4[4].data = - SIMD_FMA(grad_4[4].data, betta1_minus1_4.data, momentum_4[4].data); + momentum_4[4].data = SIMD_FMA(grad_4[4].data, betta1_minus1_4.data, momentum_4[4].data); momentum_4[5].data = SIMD_MUL(momentum_4[5].data, betta1_4.data); - momentum_4[5].data = - SIMD_FMA(grad_4[5].data, betta1_minus1_4.data, momentum_4[5].data); + momentum_4[5].data = SIMD_FMA(grad_4[5].data, betta1_minus1_4.data, momentum_4[5].data); momentum_4[6].data = SIMD_MUL(momentum_4[6].data, betta1_4.data); - momentum_4[6].data = - SIMD_FMA(grad_4[6].data, betta1_minus1_4.data, momentum_4[6].data); + momentum_4[6].data = SIMD_FMA(grad_4[6].data, betta1_minus1_4.data, momentum_4[6].data); momentum_4[7].data = SIMD_MUL(momentum_4[7].data, betta1_4.data); - momentum_4[7].data = - SIMD_FMA(grad_4[7].data, betta1_minus1_4.data, momentum_4[7].data); + momentum_4[7].data = SIMD_FMA(grad_4[7].data, betta1_minus1_4.data, momentum_4[7].data); variance_4[0].data = SIMD_MUL(variance_4[0].data, betta2_4.data); variance_4[1].data = SIMD_MUL(variance_4[1].data, betta2_4.data); @@ -491,22 +463,14 @@ void Adam_Optimizer::Step_8(float* _params, grad_4[5].data = SIMD_MUL(grad_4[5].data, grad_4[5].data); grad_4[6].data = SIMD_MUL(grad_4[6].data, grad_4[6].data); grad_4[7].data = SIMD_MUL(grad_4[7].data, grad_4[7].data); - variance_4[0].data = - SIMD_FMA(grad_4[0].data, betta2_minus1_4.data, variance_4[0].data); - variance_4[1].data = - SIMD_FMA(grad_4[1].data, betta2_minus1_4.data, variance_4[1].data); - variance_4[2].data = - SIMD_FMA(grad_4[2].data, betta2_minus1_4.data, variance_4[2].data); - variance_4[3].data = - SIMD_FMA(grad_4[3].data, betta2_minus1_4.data, variance_4[3].data); - variance_4[4].data = - SIMD_FMA(grad_4[4].data, betta2_minus1_4.data, variance_4[4].data); - variance_4[5].data = - SIMD_FMA(grad_4[5].data, betta2_minus1_4.data, variance_4[5].data); - variance_4[6].data = - SIMD_FMA(grad_4[6].data, betta2_minus1_4.data, variance_4[6].data); - variance_4[7].data = - SIMD_FMA(grad_4[7].data, betta2_minus1_4.data, variance_4[7].data); + variance_4[0].data = SIMD_FMA(grad_4[0].data, betta2_minus1_4.data, variance_4[0].data); + variance_4[1].data = SIMD_FMA(grad_4[1].data, betta2_minus1_4.data, variance_4[1].data); + variance_4[2].data = SIMD_FMA(grad_4[2].data, betta2_minus1_4.data, variance_4[2].data); + variance_4[3].data = SIMD_FMA(grad_4[3].data, betta2_minus1_4.data, variance_4[3].data); + variance_4[4].data = SIMD_FMA(grad_4[4].data, betta2_minus1_4.data, variance_4[4].data); + variance_4[5].data = SIMD_FMA(grad_4[5].data, betta2_minus1_4.data, variance_4[5].data); + variance_4[6].data = SIMD_FMA(grad_4[6].data, betta2_minus1_4.data, variance_4[6].data); + variance_4[7].data = SIMD_FMA(grad_4[7].data, betta2_minus1_4.data, variance_4[7].data); grad_4[0].data = SIMD_SQRT(variance_4[0].data); grad_4[1].data = SIMD_SQRT(variance_4[1].data); @@ -554,20 +518,15 @@ void Adam_Optimizer::Step_8(float* _params, if (dev_params) { SIMD_STORE(_doubled_buffer[_buf_index] + (i - t), param_4[0].data); - SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH, - param_4[1].data); + SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH, param_4[1].data); SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + (SIMD_WIDTH << 1), - param_4[2].data); - SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 3, - param_4[3].data); + param_4[2].data); + SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 3, param_4[3].data); SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + (SIMD_WIDTH << 2), - param_4[4].data); - SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 5, - param_4[5].data); - SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 6, - param_4[6].data); - SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 7, - param_4[7].data); + param_4[4].data); + SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 5, param_4[5].data); + SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 6, param_4[6].data); + SIMD_STORE(_doubled_buffer[_buf_index] + (i - t) + SIMD_WIDTH * 7, param_4[7].data); } SIMD_STORE(_exp_avg + i, momentum_4[0].data); diff --git a/csrc/includes/cpu_adam.h b/csrc/includes/cpu_adam.h old mode 100755 new mode 100644 index da1f73f3e44f..996450c56ae9 --- a/csrc/includes/cpu_adam.h +++ b/csrc/includes/cpu_adam.h @@ -1,15 +1,15 @@ #pragma once +#include #include #include +#include #include #include #include "context.h" #include "cublas_v2.h" #include "cuda.h" #include "curand.h" -#include -#include #define CUDA_CHECK(callstr) \ { \ @@ -23,25 +23,25 @@ #define TILE (1024 * 1024 * 1024) #if defined(__AVX512__) - #define SIMD_STORE(a, d) _mm512_storeu_ps(a, d) - #define SIMD_LOAD(x) _mm512_loadu_ps(x) - #define SIMD_SET(x) _mm512_set1_ps(x) - #define SIMD_MUL(x, y) _mm512_mul_ps(x, y) - #define SIMD_FMA(x, y, c) _mm512_fmadd_ps(x, y, c) - #define SIMD_SQRT(x) _mm512_sqrt_ps(x) - #define SIMD_DIV(x, y) _mm512_div_ps(x, y) - #define SIMD_WIDTH 16 -#else - #if defined(__AVX256__) - #define SIMD_STORE(a, d) _mm256_storeu_ps(a, d) - #define SIMD_LOAD(x) _mm256_loadu_ps(x) - #define SIMD_SET(x) _mm256_set1_ps(x) - #define SIMD_MUL(x, y) _mm256_mul_ps(x, y) - #define SIMD_FMA(x, y, c) _mm256_fmadd_ps(x, y, c) - #define SIMD_SQRT(x) _mm256_sqrt_ps(x) - #define SIMD_DIV(x, y) _mm256_div_ps(x, y) - #define SIMD_WIDTH 8 - #endif +#define SIMD_STORE(a, d) _mm512_storeu_ps(a, d) +#define SIMD_LOAD(x) _mm512_loadu_ps(x) +#define SIMD_SET(x) _mm512_set1_ps(x) +#define SIMD_MUL(x, y) _mm512_mul_ps(x, y) +#define SIMD_FMA(x, y, c) _mm512_fmadd_ps(x, y, c) +#define SIMD_SQRT(x) _mm512_sqrt_ps(x) +#define SIMD_DIV(x, y) _mm512_div_ps(x, y) +#define SIMD_WIDTH 16 +#else +#if defined(__AVX256__) +#define SIMD_STORE(a, d) _mm256_storeu_ps(a, d) +#define SIMD_LOAD(x) _mm256_loadu_ps(x) +#define SIMD_SET(x) _mm256_set1_ps(x) +#define SIMD_MUL(x, y) _mm256_mul_ps(x, y) +#define SIMD_FMA(x, y, c) _mm256_fmadd_ps(x, y, c) +#define SIMD_SQRT(x) _mm256_sqrt_ps(x) +#define SIMD_DIV(x, y) _mm256_div_ps(x, y) +#define SIMD_WIDTH 8 +#endif #endif class Adam_Optimizer { @@ -89,14 +89,15 @@ class Adam_Optimizer { inline void IncrementStep() { _betta1_t *= _betta1; - _betta2_t *= _betta2; + _betta2_t *= _betta2; } + private: #if defined(__AVX512__) or defined(__AVX256__) union AVX_Data { #if defined(__AVX512__) __m512 data; -#else +#else __m256 data; #endif // float data_f[16]; diff --git a/setup.py b/setup.py index f59e71adc2c7..f2dc2ba3b37a 100755 --- a/setup.py +++ b/setup.py @@ -106,7 +106,6 @@ def fetch_requirements(path): SIMD_WIDTH = '-D__AVX256__' print("SIMD_WIDTH = ", SIMD_WIDTH) - ext_modules = [] ## Lamb ## From 10296198587b6281e2826ee1b684a9abc29d094b Mon Sep 17 00:00:00 2001 From: Reza Yazdani Date: Tue, 8 Sep 2020 21:42:40 +0000 Subject: [PATCH 3/8] adding cpufeature to requirements --- requirements/requirements.txt | 1 + 1 file changed, 1 insertion(+) diff --git a/requirements/requirements.txt b/requirements/requirements.txt index 63f0022b314d..d9881f4bc580 100644 --- a/requirements/requirements.txt +++ b/requirements/requirements.txt @@ -2,4 +2,5 @@ torch>=1.2 torchvision>=0.4.0 tqdm psutil +cpufeature tensorboardX==1.8 From b725aba9593acaf067c4ea64cfe93582088d8177 Mon Sep 17 00:00:00 2001 From: Jeff Rasley Date: Tue, 8 Sep 2020 15:03:14 -0700 Subject: [PATCH 4/8] Update install.sh --- install.sh | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/install.sh b/install.sh index 3dae80b4033b..6c98120484a1 100755 --- a/install.sh +++ b/install.sh @@ -164,10 +164,10 @@ if [ ! -f $hostfile ]; then local_only=1 fi -#if [ "$skip_requirements" == "0" ]; then -# # Ensure dependencies are installed locally -# $PIP_SUDO $PIP_INSTALL -r requirements.txt -#fi +if [ "$skip_requirements" == "0" ]; then + # Ensure dependencies are installed locally + $PIP_SUDO $PIP_INSTALL -r requirements.txt +fi # Build wheels if [ "$third_party_install" == "1" ]; then @@ -220,10 +220,10 @@ else tmp_wheel_path="/tmp/deepspeed_wheels" pdsh -w $hosts "if [ -d $tmp_wheel_path ]; then rm $tmp_wheel_path/*.whl; else mkdir -pv $tmp_wheel_path; fi" - #pdcp -w $hosts requirements/*.txt ${tmp_wheel_path}/ - #if [ "$skip_requirements" == "0" ]; then - # pdsh -w $hosts "$PIP_SUDO $PIP_INSTALL -r ${tmp_wheel_path}/requirements.txt" - #fi + pdcp -w $hosts requirements/requirements.txt ${tmp_wheel_path}/ + if [ "$skip_requirements" == "0" ]; then + pdsh -w $hosts "$PIP_SUDO $PIP_INSTALL -r ${tmp_wheel_path}/requirements.txt" + fi if [ "$third_party_install" == "1" ]; then pdsh -w $hosts "$PIP_SUDO pip uninstall -y apex" pdcp -w $hosts third_party/apex/dist/apex*.whl $tmp_wheel_path/ From e7fe1a8cf1f52ba11468247e3d1ede8ceb5785b3 Mon Sep 17 00:00:00 2001 From: Jeff Rasley Date: Tue, 8 Sep 2020 15:11:32 -0700 Subject: [PATCH 5/8] Update install.sh --- install.sh | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/install.sh b/install.sh index 6c98120484a1..a4ee6bbf3b72 100755 --- a/install.sh +++ b/install.sh @@ -166,7 +166,7 @@ fi if [ "$skip_requirements" == "0" ]; then # Ensure dependencies are installed locally - $PIP_SUDO $PIP_INSTALL -r requirements.txt + $PIP_SUDO $PIP_INSTALL -r requirements/requirements.txt fi # Build wheels From 637cbaea2789b770e9f536cbf65a038415fa09ce Mon Sep 17 00:00:00 2001 From: Reza Yazdani Date: Tue, 8 Sep 2020 22:23:28 +0000 Subject: [PATCH 6/8] include cpu-adam in the features --- docs/_pages/features.md | 6 ++++++ docs/index.md | 1 + 2 files changed, 7 insertions(+) diff --git a/docs/_pages/features.md b/docs/_pages/features.md index 451e3b2af534..d1871604d3ce 100755 --- a/docs/_pages/features.md +++ b/docs/_pages/features.md @@ -162,6 +162,12 @@ Please see the [core API doc](https://deepspeed.readthedocs.io/) for more detail With DeepSpeed, the user can choose to use a high performance implementation of ADAM from NVIDIA, or any training optimizer that extends torch's `torch.optim.Optimizer` class. +### CPU-adam: High-Performance Vectorized implementation of Adam +We introduce an efficient implementation of Adam optimizer on CPU that improves the parameter-update +performance by nearly an order of magnitude. Comparing to torch-adam, we observe 5.1x to 6.5x +speedups considering the model-size betweein 1 to 10 billion parameters. For the CPU-Adam implementation, +we use the AVX SIMD instructions on Intel-X86 architecture. Moreover, we support both AVX-512 and AVX-2. + ### Memory bandwidth optimized FP16 Optimizer Mixed precision training is handled by the DeepSpeed FP16 Optimizer. This optimizer not only handles FP16 training but is also highly efficient. The performance of weight update diff --git a/docs/index.md b/docs/index.md index 6dea83db268f..c17533628471 100755 --- a/docs/index.md +++ b/docs/index.md @@ -167,6 +167,7 @@ overview](/features/) for descriptions and usage. * Automatic loss scaling with mixed precision * [Training Optimizers](/features/#training-optimizers) * Fused Adam optimizer and arbitrary `torch.optim.Optimizer` + * CPU-adam: High-Performance vectorized Adam * Memory bandwidth optimized FP16 Optimizer * Large Batch Training with LAMB Optimizer * Memory efficient Training with ZeRO Optimizer From b9e91f3038269d4db9362dd8efc26fa2140bbf59 Mon Sep 17 00:00:00 2001 From: Reza Yazdani Date: Tue, 8 Sep 2020 22:26:41 +0000 Subject: [PATCH 7/8] update features --- docs/_pages/features.md | 4 ++-- docs/index.md | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/docs/_pages/features.md b/docs/_pages/features.md index d1871604d3ce..c2f8ea0c254b 100755 --- a/docs/_pages/features.md +++ b/docs/_pages/features.md @@ -162,11 +162,11 @@ Please see the [core API doc](https://deepspeed.readthedocs.io/) for more detail With DeepSpeed, the user can choose to use a high performance implementation of ADAM from NVIDIA, or any training optimizer that extends torch's `torch.optim.Optimizer` class. -### CPU-adam: High-Performance Vectorized implementation of Adam +### CPU-Adam: High-Performance vectorized implementation of Adam We introduce an efficient implementation of Adam optimizer on CPU that improves the parameter-update performance by nearly an order of magnitude. Comparing to torch-adam, we observe 5.1x to 6.5x speedups considering the model-size betweein 1 to 10 billion parameters. For the CPU-Adam implementation, -we use the AVX SIMD instructions on Intel-X86 architecture. Moreover, we support both AVX-512 and AVX-2. +we use the AVX SIMD instructions on Intel-x86 architecture. We support both AVX-512 and AVX-2 instruction sets. ### Memory bandwidth optimized FP16 Optimizer Mixed precision training is handled by the DeepSpeed FP16 Optimizer. This optimizer not diff --git a/docs/index.md b/docs/index.md index c17533628471..4aa29673ecc6 100755 --- a/docs/index.md +++ b/docs/index.md @@ -167,7 +167,7 @@ overview](/features/) for descriptions and usage. * Automatic loss scaling with mixed precision * [Training Optimizers](/features/#training-optimizers) * Fused Adam optimizer and arbitrary `torch.optim.Optimizer` - * CPU-adam: High-Performance vectorized Adam + * CPU-Adam: High-Performance vectorized Adam * Memory bandwidth optimized FP16 Optimizer * Large Batch Training with LAMB Optimizer * Memory efficient Training with ZeRO Optimizer From d93f52f8c15421bf40cc92987a616efa631fed96 Mon Sep 17 00:00:00 2001 From: Reza Yazdani Date: Tue, 8 Sep 2020 22:50:51 +0000 Subject: [PATCH 8/8] update features --- docs/_pages/features.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/_pages/features.md b/docs/_pages/features.md index c2f8ea0c254b..f93f9465882b 100755 --- a/docs/_pages/features.md +++ b/docs/_pages/features.md @@ -165,7 +165,7 @@ NVIDIA, or any training optimizer that extends torch's `torch.optim.Optimizer` c ### CPU-Adam: High-Performance vectorized implementation of Adam We introduce an efficient implementation of Adam optimizer on CPU that improves the parameter-update performance by nearly an order of magnitude. Comparing to torch-adam, we observe 5.1x to 6.5x -speedups considering the model-size betweein 1 to 10 billion parameters. For the CPU-Adam implementation, +speedups considering the model-size between 1 to 10 billion parameters. For the CPU-Adam implementation, we use the AVX SIMD instructions on Intel-x86 architecture. We support both AVX-512 and AVX-2 instruction sets. ### Memory bandwidth optimized FP16 Optimizer