CUDA: use shared mem for ssm_conv (#20128)
* CUDA: use shared mem for ssm_conv * fuse silu + ssm_conv * fuse unary + mul * enable for fp16 * formatting Co-authored-by: Johannes Gäßler <johannesg@5d6.de> --------- Co-authored-by: Johannes Gäßler <johannesg@5d6.de>
This commit is contained in:
co-authored by
Johannes Gäßler
parent
388baabc06
commit
1e38a7a6fa
@@ -560,3 +560,58 @@ void ggml_cuda_op_leaky_relu(ggml_backend_cuda_context & ctx, ggml_tensor * dst)
|
||||
leaky_relu_cuda((const float *)src0_d, (float *)dst_d, ggml_nelements(src0), negative_slope, stream);
|
||||
}
|
||||
}
|
||||
|
||||
/* fused unary + mul */
|
||||
|
||||
template <float (*op)(float)>
|
||||
static void ggml_cuda_op_unary_mul_impl(ggml_backend_cuda_context & ctx, ggml_tensor * unary_node, ggml_tensor * mul_node) {
|
||||
// unary_node: UNARY op applied to unary_node->src[0]
|
||||
// mul_node: MUL(a, b) where one of a/b is unary_node
|
||||
// Output goes to mul_node->data
|
||||
|
||||
const ggml_tensor * unary_src = unary_node->src[0]; // input to the unary op
|
||||
const ggml_tensor * other_src = (mul_node->src[0] == unary_node) ? mul_node->src[1] : mul_node->src[0];
|
||||
|
||||
GGML_ASSERT(ggml_is_contiguous_1(unary_src));
|
||||
GGML_ASSERT(unary_src->nb[0] == ggml_element_size(unary_src));
|
||||
GGML_ASSERT(ggml_is_contiguous_1(other_src));
|
||||
GGML_ASSERT(other_src->nb[0] == ggml_element_size(other_src));
|
||||
GGML_ASSERT(ggml_are_same_shape(unary_src, other_src));
|
||||
|
||||
GGML_ASSERT(unary_src->type == GGML_TYPE_F32 || unary_src->type == GGML_TYPE_F16);
|
||||
GGML_ASSERT(unary_src->type == other_src->type);
|
||||
GGML_ASSERT(unary_src->type == mul_node->type);
|
||||
|
||||
cudaStream_t stream = ctx.stream();
|
||||
|
||||
const int64_t k = ggml_nelements(mul_node);
|
||||
const int64_t nc = unary_src->ne[0];
|
||||
const int64_t unary_stride = unary_src->nb[1];
|
||||
const int64_t other_stride = other_src->nb[1];
|
||||
|
||||
if (unary_src->type == GGML_TYPE_F16) {
|
||||
unary_gated_cuda<op>((const half *) unary_src->data, (const half *) other_src->data,
|
||||
(half *) mul_node->data, k, nc,
|
||||
unary_stride / sizeof(half), other_stride / sizeof(half), stream);
|
||||
} else {
|
||||
unary_gated_cuda<op>((const float *) unary_src->data, (const float *) other_src->data,
|
||||
(float *) mul_node->data, k, nc,
|
||||
unary_stride / sizeof(float), other_stride / sizeof(float), stream);
|
||||
}
|
||||
}
|
||||
|
||||
void ggml_cuda_op_unary_mul(ggml_backend_cuda_context & ctx, ggml_tensor * unary_node, ggml_tensor * mul_node) {
|
||||
switch (ggml_get_unary_op(unary_node)) {
|
||||
case GGML_UNARY_OP_SILU:
|
||||
ggml_cuda_op_unary_mul_impl<op_silu>(ctx, unary_node, mul_node);
|
||||
break;
|
||||
case GGML_UNARY_OP_SIGMOID:
|
||||
ggml_cuda_op_unary_mul_impl<op_sigmoid>(ctx, unary_node, mul_node);
|
||||
break;
|
||||
case GGML_UNARY_OP_SOFTPLUS:
|
||||
ggml_cuda_op_unary_mul_impl<op_softplus>(ctx, unary_node, mul_node);
|
||||
break;
|
||||
default:
|
||||
GGML_ABORT("Unsupported unary op for fused unary+mul");
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user