CUDA: fuse rms -> add -> rms (#2297)
This commit is contained in:
parent
3c949f3399
commit
8b276c08ef
|
|
@ -3939,9 +3939,49 @@ 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 {
|
||||
}
|
||||
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:
|
||||
ggml_cuda_op_fused_rms_rms_add(ctx, dst);
|
||||
|
|
|
|||
|
|
@ -1053,3 +1053,183 @@ void ggml_cuda_op_fused_rms_rms_add(ggml_backend_cuda_context & ctx, ggml_tensor
|
|||
GGML_ABORT("Not implemented");
|
||||
}
|
||||
}
|
||||
|
||||
template <int block_size, typename src_t>
|
||||
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 <typename src_t>
|
||||
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><<<nrows, block_dims, 0, stream>>>(ncols, dst1, dst2, x, c1, c2, a, eps1, eps2);
|
||||
} else {
|
||||
const dim3 block_dims(1024, 1, 1);
|
||||
fused_rms_add_rms_f32<1024><<<nrows, block_dims, 0, stream>>>(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 <int block_size, typename src_t>
|
||||
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 <typename src_t>
|
||||
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><<<nrows, block_dims, 0, stream>>>(ncols, dst, x, c, a, eps);
|
||||
} else {
|
||||
const dim3 block_dims(1024, 1, 1);
|
||||
fused_rms_add_f32<1024><<<nrows, block_dims, 0, stream>>>(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());
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
Loading…
Reference in New Issue