| 703 | |
| 704 | template <typename T, typename Tindex> |
| 705 | __global__ __launch_bounds__(1024) void SparseApplyAdamKernel( |
| 706 | T* var, T* m, T* v, const T* grad, const T* beta1_power, const T* beta2_power, |
| 707 | const T* lr, const T* beta1, const T* beta2, const T* epsilon, const Tindex* indices, |
| 708 | Tindex param_rows, Tindex updates_size, Tindex indices_size) { |
| 709 | Tindex col_size = updates_size / indices_size; |
| 710 | const T alpha = (*lr) * sqrt(static_cast<T>(1) - *beta2_power) / |
| 711 | (static_cast<T>(1) - *beta1_power); |
| 712 | |
| 713 | GPU_1D_KERNEL_LOOP(grad_index, updates_size) { |
| 714 | Tindex indices_row = grad_index / col_size; |
| 715 | Tindex param_row = indices[indices_row]; |
| 716 | if (param_row < 0 || param_row >= param_rows) { |
| 717 | // Ignore indices that are out of range |
| 718 | continue; |
| 719 | } |
| 720 | |
| 721 | // Index of var, m and v |
| 722 | Tindex param_index = param_row*col_size + grad_index%col_size; |
| 723 | const T& g = grad[grad_index]; |
| 724 | T& var_a = var[param_index]; |
| 725 | T& m_a = m[param_index]; |
| 726 | T& v_a = v[param_index]; |
| 727 | |
| 728 | m_a += (g-m_a) * (static_cast<T>(1) - (*beta1)); |
| 729 | v_a += (g*g - v_a) * (static_cast<T>(1) - (*beta2)); |
| 730 | var_a -= (m_a*alpha) / (sqrt(v_a)+ (*epsilon)); |
| 731 | } |
| 732 | } |
| 733 | template <typename T, typename Tindex> |
| 734 | struct SparseApplyAdam<GPUDevice, T, Tindex> { |
| 735 | Status operator()(const GPUDevice& d, typename TTypes<T>::Matrix var, |