metal: add col2im_1d op (f32/f16/bf16) (#25176)
* metal: add col2im_1d op (f32/f16/bf16) Gather kernel mirroring the CPU/CUDA path: each output (t_out, oc) reads its ceil(K/s0) source columns with an F32 accumulator, a single write and no atomics. One thread per output element, 256 per threadgroup. * metal: check dst contiguity and type match in supports_op for COL2IM_1D Align the GGML_OP_COL2IM_1D predicate with the CPU, CUDA, and Vulkan backends: the kernel writes dst with linear indexing and assumes the same type as src0, so supports_op must also require a contiguous dst and op->type == op->src[0]->type. * Update ggml/src/ggml-metal/ggml-metal.metal Co-authored-by: YiChen Lv <63285796+forforever73@users.noreply.github.com> --------- Co-authored-by: YiChen Lv <63285796+forforever73@users.noreply.github.com>
This commit is contained in:
@@ -1800,6 +1800,26 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_conv_transpose_1
|
|||||||
return res;
|
return res;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_col2im_1d(ggml_metal_library_t lib, const ggml_tensor * op) {
|
||||||
|
assert(op->op == GGML_OP_COL2IM_1D);
|
||||||
|
|
||||||
|
GGML_ASSERT(ggml_is_contiguous(op->src[0]));
|
||||||
|
GGML_ASSERT(op->src[0]->type == GGML_TYPE_F32 || op->src[0]->type == GGML_TYPE_F16 || op->src[0]->type == GGML_TYPE_BF16);
|
||||||
|
|
||||||
|
char base[256];
|
||||||
|
char name[256];
|
||||||
|
|
||||||
|
snprintf(base, 256, "kernel_col2im_1d_%s", ggml_type_name(op->src[0]->type));
|
||||||
|
snprintf(name, 256, "%s", base);
|
||||||
|
|
||||||
|
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
|
||||||
|
if (!res.pipeline) {
|
||||||
|
res = ggml_metal_library_compile_pipeline(lib, base, name, nullptr);
|
||||||
|
}
|
||||||
|
|
||||||
|
return res;
|
||||||
|
}
|
||||||
|
|
||||||
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_conv_transpose_2d(ggml_metal_library_t lib, const ggml_tensor * op) {
|
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_conv_transpose_2d(ggml_metal_library_t lib, const ggml_tensor * op) {
|
||||||
assert(op->op == GGML_OP_CONV_TRANSPOSE_2D);
|
assert(op->op == GGML_OP_CONV_TRANSPOSE_2D);
|
||||||
|
|
||||||
|
|||||||
@@ -150,6 +150,7 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_rope
|
|||||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_im2col (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_im2col (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_conv_transpose_1d (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_conv_transpose_1d (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_conv_transpose_2d (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_conv_transpose_2d (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||||
|
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_col2im_1d (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_conv_2d (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_conv_2d (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_conv_3d (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_conv_3d (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_upscale (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_upscale (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||||
|
|||||||
@@ -1157,6 +1157,11 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te
|
|||||||
(op->src[0]->type == GGML_TYPE_F16 || op->src[0]->type == GGML_TYPE_F32) &&
|
(op->src[0]->type == GGML_TYPE_F16 || op->src[0]->type == GGML_TYPE_F32) &&
|
||||||
op->src[1]->type == GGML_TYPE_F32 &&
|
op->src[1]->type == GGML_TYPE_F32 &&
|
||||||
op->type == GGML_TYPE_F32;
|
op->type == GGML_TYPE_F32;
|
||||||
|
case GGML_OP_COL2IM_1D:
|
||||||
|
return (op->src[0]->type == GGML_TYPE_F32 || op->src[0]->type == GGML_TYPE_F16 || op->src[0]->type == GGML_TYPE_BF16) &&
|
||||||
|
op->type == op->src[0]->type &&
|
||||||
|
ggml_is_contiguous(op->src[0]) &&
|
||||||
|
ggml_is_contiguous(op);
|
||||||
case GGML_OP_CONV_3D:
|
case GGML_OP_CONV_3D:
|
||||||
return ggml_is_contiguous(op->src[0]) &&
|
return ggml_is_contiguous(op->src[0]) &&
|
||||||
ggml_is_contiguous(op->src[1]) &&
|
ggml_is_contiguous(op->src[1]) &&
|
||||||
|
|||||||
@@ -603,6 +603,16 @@ typedef struct {
|
|||||||
uint64_t nb1;
|
uint64_t nb1;
|
||||||
} ggml_metal_kargs_conv_transpose_1d;
|
} ggml_metal_kargs_conv_transpose_1d;
|
||||||
|
|
||||||
|
typedef struct {
|
||||||
|
int32_t T_in;
|
||||||
|
int32_t T_out;
|
||||||
|
int32_t OC;
|
||||||
|
int32_t K;
|
||||||
|
int32_t K_OC;
|
||||||
|
int32_t s0;
|
||||||
|
int32_t p0;
|
||||||
|
} ggml_metal_kargs_col2im_1d;
|
||||||
|
|
||||||
typedef struct {
|
typedef struct {
|
||||||
int32_t IC;
|
int32_t IC;
|
||||||
int32_t IH;
|
int32_t IH;
|
||||||
|
|||||||
@@ -395,6 +395,10 @@ static int ggml_metal_op_encode_impl(ggml_metal_op_t ctx, int idx) {
|
|||||||
{
|
{
|
||||||
n_fuse = ggml_metal_op_conv_transpose_2d(ctx, idx);
|
n_fuse = ggml_metal_op_conv_transpose_2d(ctx, idx);
|
||||||
} break;
|
} break;
|
||||||
|
case GGML_OP_COL2IM_1D:
|
||||||
|
{
|
||||||
|
n_fuse = ggml_metal_op_col2im_1d(ctx, idx);
|
||||||
|
} break;
|
||||||
case GGML_OP_CONV_3D:
|
case GGML_OP_CONV_3D:
|
||||||
{
|
{
|
||||||
n_fuse = ggml_metal_op_conv_3d(ctx, idx);
|
n_fuse = ggml_metal_op_conv_3d(ctx, idx);
|
||||||
@@ -3854,6 +3858,47 @@ int ggml_metal_op_conv_transpose_1d(ggml_metal_op_t ctx, int idx) {
|
|||||||
return 1;
|
return 1;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
int ggml_metal_op_col2im_1d(ggml_metal_op_t ctx, int idx) {
|
||||||
|
ggml_tensor * op = ctx->node(idx);
|
||||||
|
|
||||||
|
ggml_metal_library_t lib = ctx->lib;
|
||||||
|
ggml_metal_encoder_t enc = ctx->enc;
|
||||||
|
|
||||||
|
const int32_t s0 = ((const int32_t *)(op->op_params))[0];
|
||||||
|
const int32_t OC = ((const int32_t *)(op->op_params))[1];
|
||||||
|
const int32_t p0 = ((const int32_t *)(op->op_params))[2];
|
||||||
|
|
||||||
|
const int32_t K_OC = (int32_t) op->src[0]->ne[0];
|
||||||
|
const int32_t T_in = (int32_t) op->src[0]->ne[1];
|
||||||
|
const int32_t K = K_OC / OC;
|
||||||
|
const int32_t T_out = (int32_t) op->ne[0];
|
||||||
|
|
||||||
|
ggml_metal_kargs_col2im_1d args = {
|
||||||
|
/*.T_in =*/ T_in,
|
||||||
|
/*.T_out =*/ T_out,
|
||||||
|
/*.OC =*/ OC,
|
||||||
|
/*.K =*/ K,
|
||||||
|
/*.K_OC =*/ K_OC,
|
||||||
|
/*.s0 =*/ s0,
|
||||||
|
/*.p0 =*/ p0,
|
||||||
|
};
|
||||||
|
|
||||||
|
auto pipeline = ggml_metal_library_get_pipeline_col2im_1d(lib, op);
|
||||||
|
|
||||||
|
const int total = T_out * OC;
|
||||||
|
const int nth = 256;
|
||||||
|
const int ntg = (total + nth - 1) / nth;
|
||||||
|
|
||||||
|
ggml_metal_encoder_set_pipeline(enc, pipeline);
|
||||||
|
ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), 0);
|
||||||
|
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[0]), 1);
|
||||||
|
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op), 2);
|
||||||
|
|
||||||
|
ggml_metal_encoder_dispatch_threadgroups(enc, ntg, 1, 1, nth, 1, 1);
|
||||||
|
|
||||||
|
return 1;
|
||||||
|
}
|
||||||
|
|
||||||
int ggml_metal_op_conv_transpose_2d(ggml_metal_op_t ctx, int idx) {
|
int ggml_metal_op_conv_transpose_2d(ggml_metal_op_t ctx, int idx) {
|
||||||
ggml_tensor * op = ctx->node(idx);
|
ggml_tensor * op = ctx->node(idx);
|
||||||
|
|
||||||
|
|||||||
@@ -78,6 +78,7 @@ int ggml_metal_op_conv_2d (ggml_metal_op_t ctx, int idx);
|
|||||||
int ggml_metal_op_conv_3d (ggml_metal_op_t ctx, int idx);
|
int ggml_metal_op_conv_3d (ggml_metal_op_t ctx, int idx);
|
||||||
int ggml_metal_op_conv_transpose_1d (ggml_metal_op_t ctx, int idx);
|
int ggml_metal_op_conv_transpose_1d (ggml_metal_op_t ctx, int idx);
|
||||||
int ggml_metal_op_conv_transpose_2d (ggml_metal_op_t ctx, int idx);
|
int ggml_metal_op_conv_transpose_2d (ggml_metal_op_t ctx, int idx);
|
||||||
|
int ggml_metal_op_col2im_1d (ggml_metal_op_t ctx, int idx);
|
||||||
int ggml_metal_op_upscale (ggml_metal_op_t ctx, int idx);
|
int ggml_metal_op_upscale (ggml_metal_op_t ctx, int idx);
|
||||||
int ggml_metal_op_pad (ggml_metal_op_t ctx, int idx);
|
int ggml_metal_op_pad (ggml_metal_op_t ctx, int idx);
|
||||||
int ggml_metal_op_pad_reflect_1d (ggml_metal_op_t ctx, int idx);
|
int ggml_metal_op_pad_reflect_1d (ggml_metal_op_t ctx, int idx);
|
||||||
|
|||||||
@@ -4977,6 +4977,49 @@ kernel void kernel_conv_transpose_1d<half>(
|
|||||||
uint3 tgpg[[threadgroups_per_grid]]);
|
uint3 tgpg[[threadgroups_per_grid]]);
|
||||||
|
|
||||||
|
|
||||||
|
template <typename T>
|
||||||
|
kernel void kernel_col2im_1d(
|
||||||
|
constant ggml_metal_kargs_col2im_1d & args,
|
||||||
|
device const T * col,
|
||||||
|
device T * dst,
|
||||||
|
uint tgpig [[threadgroup_position_in_grid]],
|
||||||
|
uint tpitg [[thread_position_in_threadgroup]],
|
||||||
|
uint ntg [[threads_per_threadgroup]]) {
|
||||||
|
|
||||||
|
const int idx = tgpig * ntg + tpitg;
|
||||||
|
if (idx >= args.T_out * args.OC) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
const int t_out = idx % args.T_out;
|
||||||
|
const int oc = idx / args.T_out;
|
||||||
|
const int t_abs = t_out + args.p0; // absolute position in uncropped signal
|
||||||
|
|
||||||
|
int t_in_min = (t_abs - args.K + args.s0) / args.s0; // ceil((t_abs - K + 1) / s0)
|
||||||
|
if (t_in_min < 0) {
|
||||||
|
t_in_min = 0;
|
||||||
|
}
|
||||||
|
int t_in_max = t_abs / args.s0;
|
||||||
|
if (t_in_max >= args.T_in) {
|
||||||
|
t_in_max = args.T_in - 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
float sum = 0.0f;
|
||||||
|
for (int t_in = t_in_min; t_in <= t_in_max; t_in++) {
|
||||||
|
const int k = t_abs - t_in * args.s0;
|
||||||
|
sum += float(col[(oc * args.K + k) + t_in * args.K_OC]);
|
||||||
|
}
|
||||||
|
|
||||||
|
dst[t_out + oc * args.T_out] = T(sum);
|
||||||
|
}
|
||||||
|
|
||||||
|
template [[host_name("kernel_col2im_1d_f32")]] kernel void kernel_col2im_1d<float>(constant ggml_metal_kargs_col2im_1d &, device const float *, device float *, uint, uint, uint);
|
||||||
|
template [[host_name("kernel_col2im_1d_f16")]] kernel void kernel_col2im_1d<half>(constant ggml_metal_kargs_col2im_1d &, device const half *, device half *, uint, uint, uint);
|
||||||
|
#if defined(GGML_METAL_HAS_BF16)
|
||||||
|
template [[host_name("kernel_col2im_1d_bf16")]] kernel void kernel_col2im_1d<bfloat>(constant ggml_metal_kargs_col2im_1d &, device const bfloat *, device bfloat *, uint, uint, uint);
|
||||||
|
#endif
|
||||||
|
|
||||||
|
|
||||||
typedef void (conv_transpose_2d_t)(
|
typedef void (conv_transpose_2d_t)(
|
||||||
constant ggml_metal_kargs_conv_transpose_2d & args,
|
constant ggml_metal_kargs_conv_transpose_2d & args,
|
||||||
device const float * src0,
|
device const float * src0,
|
||||||
|
|||||||
Reference in New Issue
Block a user