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) {
assert(op->op == GGML_OP_CONV_TRANSPOSE_2D);
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_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_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);
(op->src[0]->type == GGML_TYPE_F16 || op->src[0]->type == GGML_TYPE_F32) &&
op->src[1]->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:
return ggml_is_contiguous(op->src[0]) &&
ggml_is_contiguous(op->src[1]) &&
uint64_t nb1;
} 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 {
int32_t IC;
int32_t IH;
{
n_fuse = ggml_metal_op_conv_transpose_2d(ctx, idx);
} break;
+ case GGML_OP_COL2IM_1D:
+ {
+ n_fuse = ggml_metal_op_col2im_1d(ctx, idx);
+ } break;
case GGML_OP_CONV_3D:
{
n_fuse = ggml_metal_op_conv_3d(ctx, idx);
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) {
ggml_tensor * op = ctx->node(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_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_pad (ggml_metal_op_t ctx, int idx);
int ggml_metal_op_pad_reflect_1d (ggml_metal_op_t ctx, int idx);
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)(
constant ggml_metal_kargs_conv_transpose_2d & args,
device const float * src0,