diff --git a/ggml/src/ggml-cuda.cu b/ggml/src/ggml-cuda.cu index 3a6b6efd..d68ccefd 100644 --- a/ggml/src/ggml-cuda.cu +++ b/ggml/src/ggml-cuda.cu @@ -3939,8 +3939,48 @@ static bool ggml_cuda_compute_forward(ggml_backend_cuda_context & ctx, struct gg dst->ne[2] == 1 && cgraph->nodes[i+2]->ne[2] == 1) { ggml_cuda_op_fused_rms_rms_norm(ctx, dst, cgraph->nodes[i+2]); i += 2; - } else { - ggml_cuda_op_fused_rms_norm(ctx, dst); + } + else { + int inow = i; + // First try rms -> add -> rms + // This doesn't always work because the second rms result may get allocated on top of the + // first rms source + if (fusion && i + 2 < cgraph->n_nodes && + cgraph->nodes[i+1]->op == GGML_OP_ADD && + cgraph->nodes[i+2]->op == GGML_OP_FUSED_RMS_NORM && + dst->src[0]->ne[1] == 1 && + dst->src[0]->type != GGML_TYPE_Q8_0 && // In case someone has decided to use Q8_0 as the graph reduce type + cgraph->nodes[i+1]->src[0] == dst && + cgraph->nodes[i+2]->src[0] == cgraph->nodes[i+1] && + ggml_are_same_shape(dst, cgraph->nodes[i+1]->src[1])) { + auto src0 = (const char *)dst->src[0]->data; + auto src0_end = src0 + ggml_nbytes(dst->src[0]); + auto add1 = (const char *)cgraph->nodes[i+1]->data; + auto rms2 = (const char *)cgraph->nodes[i+2]->data; + auto nbytes = ggml_nbytes(dst); + bool overlap1 = add1 > src0 && add1 < src0_end; + bool overlap2 = add1 + nbytes > src0 && add1 + nbytes < src0_end; + bool overlap3 = add1 <= src0 && add1 + nbytes >= src0_end && !(add1 == src0 && add1 + nbytes == src0_end); + bool overlap4 = rms2 > src0 && rms2 < src0_end; + bool overlap5 = rms2 + nbytes > src0 && rms2 + nbytes < src0_end; + bool overlap6 = rms2 <= src0 && rms2 + nbytes >= src0_end && !(rms2 == src0 && rms2 + nbytes == src0_end); + if (!overlap1 && !overlap2 && !overlap3 && !overlap4 && !overlap5 && !overlap6) { + ggml_cuda_op_fused_rms_add_rms(ctx, cgraph->nodes[i+2]); + i += 2; + } + } + // If that did not work, try rms -> add + if (fusion && inow == i && i + 1 < cgraph->n_nodes && + dst->src[0]->type != GGML_TYPE_Q8_0 && // In case someone has decided to use Q8_0 as the graph reduce type + cgraph->nodes[i+1]->op == GGML_OP_ADD && + cgraph->nodes[i+1]->src[0] == dst && + ggml_are_same_shape(dst, cgraph->nodes[i+1]->src[1])) { + ggml_cuda_op_fused_rms_add(ctx, cgraph->nodes[i+1]); + i += 1; + } + if (inow == i) { + ggml_cuda_op_fused_rms_norm(ctx, dst); + } } break; case GGML_OP_FUSED_RMS_RMS_ADD: diff --git a/ggml/src/ggml-cuda/norm.cu b/ggml/src/ggml-cuda/norm.cu index 0a112529..9b7ae05b 100644 --- a/ggml/src/ggml-cuda/norm.cu +++ b/ggml/src/ggml-cuda/norm.cu @@ -1053,3 +1053,183 @@ void ggml_cuda_op_fused_rms_rms_add(ggml_backend_cuda_context & ctx, ggml_tensor GGML_ABORT("Not implemented"); } } + +template +static __global__ void fused_rms_add_rms_f32(int ncols, float * dst1, float * dst2, + const src_t * x, const float * c1, const float * c2, const float * a, float eps1, float eps2) { + const int row = blockIdx.x*blockDim.y + threadIdx.y; + const int tid = threadIdx.x; + + x += row*ncols; + a += row*ncols; + dst1 += row*ncols; + dst2 += row*ncols; + + float tmp = 0.0f; + for (int col = tid; col < ncols; col += block_size) { + float xc = (float)x[col]; + tmp += xc*xc; + } + + __shared__ float sum[block_size/WARP_SIZE]; + + int warp_id = tid / WARP_SIZE; + int lane_id = tid % WARP_SIZE; + + tmp = warp_reduce_sum(tmp); + if (lane_id == 0) sum[warp_id] = tmp; + __syncthreads(); + tmp = lane_id < block_size/WARP_SIZE ? sum[lane_id] : 0.0f; + tmp = warp_reduce_sum(tmp); + + float scale1 = rsqrtf(tmp/ncols + eps1); + + tmp = 0.0f; + for (int col = tid; col < ncols; col += block_size) { + float y = scale1 * c1[col] * (float)x[col] + a[col]; + //dst1[col] = y; + tmp += y*y; + } + + tmp = warp_reduce_sum(tmp); + if (lane_id == 0) sum[warp_id] = tmp; + __syncthreads(); + tmp = lane_id < block_size/WARP_SIZE ? sum[lane_id] : 0.0f; + tmp = warp_reduce_sum(tmp); + + float scale2 = rsqrtf(tmp/ncols + eps2); + + for (int col = tid; col < ncols; col += block_size) { + float y = scale1 * c1[col] * (float)x[col] + a[col]; + dst2[col] = scale2 * c2[col] * y; + dst1[col] = y; + } +} + +template +static void fused_rms_add_rms_f32_cuda(int ncols, int nrows, float * dst1, float * dst2, + const src_t * x, const float * c1, const float * c2, const float * a, + float eps1, float eps2, cudaStream_t stream) { + if (ncols < 1024) { + const dim3 block_dims(256, 1, 1); + fused_rms_add_rms_f32<256><<>>(ncols, dst1, dst2, x, c1, c2, a, eps1, eps2); + } else { + const dim3 block_dims(1024, 1, 1); + fused_rms_add_rms_f32<1024><<>>(ncols, dst1, dst2, x, c1, c2, a, eps1, eps2); + } +} + +void ggml_cuda_op_fused_rms_add_rms(ggml_backend_cuda_context & ctx, ggml_tensor * rms2) { + GGML_ASSERT(ggml_is_contiguous(rms2->src[0]) && ggml_is_contiguous(rms2)); + GGML_ASSERT(ggml_are_same_shape(rms2->src[0], rms2)); + auto add1 = rms2->src[0]; + GGML_ASSERT(add1->op == GGML_OP_ADD); + GGML_ASSERT(ggml_are_same_shape(add1->src[0], add1->src[1])); + GGML_ASSERT(ggml_is_contiguous(add1->src[0]) && ggml_is_contiguous(add1->src[1])); + GGML_ASSERT(rms2->type == GGML_TYPE_F32 && rms2->src[1]->type == GGML_TYPE_F32); + GGML_ASSERT(add1->type == GGML_TYPE_F32 && add1->src[1]->type == GGML_TYPE_F32); + auto rms1 = add1->src[0]; + GGML_ASSERT(rms1->op == GGML_OP_FUSED_RMS_NORM); + auto src0 = rms1->src[0]; + GGML_ASSERT(ggml_is_contiguous(src0) && ggml_is_contiguous(rms1)); + GGML_ASSERT(ggml_are_same_shape(src0, rms1)); + GGML_ASSERT(src0->type == GGML_TYPE_F16 || src0->type == GGML_TYPE_BF16 || src0->type == GGML_TYPE_F32); + + float eps1, eps2; + memcpy(&eps1, rms1->op_params, sizeof(float)); + memcpy(&eps2, rms2->op_params, sizeof(float)); + + if (src0->type == GGML_TYPE_F32) { + fused_rms_add_rms_f32_cuda(rms2->ne[0], ggml_nrows(rms2), (float *)add1->data, (float *)rms2->data, + (const float *)src0->data, (const float *)rms1->src[1]->data, + (const float *)rms2->src[1]->data, (const float *)add1->src[1]->data, eps1, eps2, ctx.stream()); + } + else if (src0->type == GGML_TYPE_F16) { + fused_rms_add_rms_f32_cuda(rms2->ne[0], ggml_nrows(rms2), (float *)add1->data, (float *)rms2->data, + (const half *)src0->data, (const float *)rms1->src[1]->data, + (const float *)rms2->src[1]->data, (const float *)add1->src[1]->data, eps1, eps2, ctx.stream()); + } + else { + fused_rms_add_rms_f32_cuda(rms2->ne[0], ggml_nrows(rms2), (float *)add1->data, (float *)rms2->data, + (const nv_bfloat16 *)src0->data, (const float *)rms1->src[1]->data, + (const float *)rms2->src[1]->data, (const float *)add1->src[1]->data, eps1, eps2, ctx.stream()); + } +} + +template +static __global__ void fused_rms_add_f32(int ncols, float * dst, + const src_t * x, const float * c, const float * a, float eps) { + const int row = blockIdx.x*blockDim.y + threadIdx.y; + const int tid = threadIdx.x; + + x += row*ncols; + a += row*ncols; + dst += row*ncols; + + float tmp = 0.0f; + for (int col = tid; col < ncols; col += block_size) { + float xc = (float)x[col]; + tmp += xc*xc; + } + + __shared__ float sum[block_size/WARP_SIZE]; + + int warp_id = tid / WARP_SIZE; + int lane_id = tid % WARP_SIZE; + + tmp = warp_reduce_sum(tmp); + if (lane_id == 0) sum[warp_id] = tmp; + __syncthreads(); + tmp = lane_id < block_size/WARP_SIZE ? sum[lane_id] : 0.0f; + tmp = warp_reduce_sum(tmp); + + float scale = rsqrtf(tmp/ncols + eps); + + for (int col = tid; col < ncols; col += block_size) { + dst[col] = scale * c[col] * (float)x[col] + a[col]; + } +} + +template +static void fused_rms_add_f32_cuda(int ncols, int nrows, float * dst, + const src_t * x, const float * c, const float * a, + float eps, cudaStream_t stream) { + if (ncols < 1024) { + const dim3 block_dims(256, 1, 1); + fused_rms_add_f32<256><<>>(ncols, dst, x, c, a, eps); + } else { + const dim3 block_dims(1024, 1, 1); + fused_rms_add_f32<1024><<>>(ncols, dst, x, c, a, eps); + } +} + +void ggml_cuda_op_fused_rms_add(ggml_backend_cuda_context & ctx, ggml_tensor * add) { + GGML_ASSERT(ggml_are_same_shape(add->src[0], add->src[1])); + GGML_ASSERT(add->op == GGML_OP_ADD); + auto rms = add->src[0]; + auto src = rms->src[0]; + GGML_ASSERT(ggml_is_contiguous(src)); + GGML_ASSERT(ggml_are_same_shape(src, rms)); + GGML_ASSERT(ggml_are_same_shape(add, rms)); + GGML_ASSERT(src->type == GGML_TYPE_F16 || src->type == GGML_TYPE_F32 || src->type == GGML_TYPE_BF16); + GGML_ASSERT(rms->src[1]->type == GGML_TYPE_F32); + GGML_ASSERT(add->src[1]->type == GGML_TYPE_F32); + GGML_ASSERT(ggml_nrows(rms->src[1]) == 1); + + float eps; + memcpy(&eps, rms->op_params, sizeof(float)); + + if (src->type == GGML_TYPE_F32) { + fused_rms_add_f32_cuda(rms->ne[0], ggml_nrows(rms), (float *)add->data, + (const float *)src->data, (const float *)rms->src[1]->data, + (const float *)add->src[1]->data, eps, ctx.stream()); + } else if (src->type == GGML_TYPE_F16) { + fused_rms_add_f32_cuda(rms->ne[0], ggml_nrows(rms), (float *)add->data, + (const half *)src->data, (const float *)rms->src[1]->data, + (const float *)add->src[1]->data, eps, ctx.stream()); + } else { + fused_rms_add_f32_cuda(rms->ne[0], ggml_nrows(rms), (float *)add->data, + (const nv_bfloat16 *)src->data, (const float *)rms->src[1]->data, + (const float *)add->src[1]->data, eps, ctx.stream()); + } +} diff --git a/ggml/src/ggml-cuda/norm.cuh b/ggml/src/ggml-cuda/norm.cuh index 513af1e2..0b498606 100644 --- a/ggml/src/ggml-cuda/norm.cuh +++ b/ggml/src/ggml-cuda/norm.cuh @@ -17,3 +17,7 @@ void ggml_cuda_op_fused_add_add_rms_norm(ggml_backend_cuda_context & ctx, ggml_t void ggml_cuda_op_fused_rms_rms_norm(ggml_backend_cuda_context & ctx, ggml_tensor * rms1, ggml_tensor * rms2); void ggml_cuda_op_fused_rms_rms_add(ggml_backend_cuda_context & ctx, ggml_tensor * dst); + +void ggml_cuda_op_fused_rms_add_rms(ggml_backend_cuda_context & ctx, ggml_tensor * rms2); + +void ggml_cuda_op_fused_rms_add(ggml_backend_cuda_context & ctx, ggml_tensor * dst);