From aa328c2c20f685be5f5573198f56f5c9bbc7c862 Mon Sep 17 00:00:00 2001 From: Sherlock Date: Fri, 24 Jul 2020 10:54:31 -0700 Subject: [PATCH] Update GratherGard to accumulate in fp32 (#4601) --- .../training_ops/cuda/tensor/gather_grad_impl.cu | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/orttraining/orttraining/training_ops/cuda/tensor/gather_grad_impl.cu b/orttraining/orttraining/training_ops/cuda/tensor/gather_grad_impl.cu index 1fc15f2218..d55a6bae4e 100644 --- a/orttraining/orttraining/training_ops/cuda/tensor/gather_grad_impl.cu +++ b/orttraining/orttraining/training_ops/cuda/tensor/gather_grad_impl.cu @@ -40,15 +40,15 @@ __global__ void _GatherGradImpl( const int weight_row = itr * input_numel + ((int)input[idx]) * stride; //the offset of the input const int grad_row = (itr * numel + ((int)indices[idx])) * stride; //the offset of the gradient - T gradient[SZ]; - T weight[SZ]; + float gradient[SZ]; + float weight[SZ]; #pragma unroll for (int ii = 0; ii < SZ; ii++) { int feature_dim = start_feature + ii * GPU_WARP_SIZE; if (feature_dim < stride) { - gradient[ii] = static_cast(grad_output[grad_row + feature_dim]); - weight[ii] = static_cast(grad_weight[weight_row + feature_dim]); + gradient[ii] = static_cast(grad_output[grad_row + feature_dim]); + weight[ii] = static_cast(grad_weight[weight_row + feature_dim]); } }