Commit 4364bf723 for llama.cpp
commit 4364bf7232e65c34eca8d9500c5464389662de6b
Author: Pascal <admin@serveurperso.com>
Date: Mon Sep 28 12:26:50 2026 +0200
metal: support left and circular padding in GGML_OP_PAD (#29561)
* metal: support left and circular padding in GGML_OP_PAD
Align Metal with CPU, CUDA and Vulkan: shift the source coordinates by
the left paddings, wrap them around with the same wrap_around when
circular, and read the source through nb00, which also fixes a right
padding of a permuted source. A test case covers it.
Drop the f32_4 kernel: its selection is disabled as slower, and it
fails two pad cases once enabled.
* metal: use a function constant for the circular pad variant
Address review from ggerganov: replace the bool template with FC_PAD,
as FC_upscale_aa does, so the pad kernel is compiled once and
specialized per pipeline.
diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp
index dc6b695eb..95b6c513f 100644
--- a/ggml/src/ggml-metal/ggml-metal-device.cpp
+++ b/ggml/src/ggml-metal/ggml-metal-device.cpp
@@ -2397,21 +2397,21 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_pad(ggml_metal_l
char base[256];
char name[256];
- // note: this is slower
- //const bool is_c4 = op->src[0]->ne[0] % 4 == 0 && op->ne[0] % 4 == 0;
- const bool is_c4 = false;
+ const bool circular = ggml_get_op_params_i32(op, 8) != 0;
- snprintf(base, 256, "kernel_pad_%s%s", ggml_type_name(op->src[0]->type), is_c4 ? "_4" : "");
- snprintf(name, 256, "%s", base);
+ snprintf(base, 256, "kernel_pad_%s", ggml_type_name(op->src[0]->type));
+ snprintf(name, 256, "%s_circular=%d", base, circular);
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
- if (res.pipeline) {
- return res;
- }
+ if (!res.pipeline) {
+ ggml_metal_cv_t cv = ggml_metal_cv_init();
- res = ggml_metal_library_compile_pipeline(lib, base, name, nullptr);
+ ggml_metal_cv_set_bool(cv, circular, FC_PAD + 0);
- res.c4 = is_c4;
+ res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
+
+ ggml_metal_cv_free(cv);
+ }
return res;
}
diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m
index fa58b8965..9c07ff6e5 100644
--- a/ggml/src/ggml-metal/ggml-metal-device.m
+++ b/ggml/src/ggml-metal/ggml-metal-device.m
@@ -1713,13 +1713,6 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te
case GGML_OP_POOL_2D:
return op->src[0]->type == GGML_TYPE_F32;
case GGML_OP_PAD:
- // TODO: add circular padding support for metal, see https://github.com/ggml-org/llama.cpp/pull/16985
- if (ggml_get_op_params_i32(op, 8) != 0) {
- return false;
- }
-
- return (ggml_get_op_params_i32(op, 0) == 0) && (ggml_get_op_params_i32(op, 2) == 0) &&
- (ggml_get_op_params_i32(op, 4) == 0) && (ggml_get_op_params_i32(op, 6) == 0);
case GGML_OP_PAD_REFLECT_1D:
case GGML_OP_TIMESTEP_EMBEDDING:
return op->src[0]->type == GGML_TYPE_F32;
diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h
index eed85f283..8bb70d077 100644
--- a/ggml/src/ggml-metal/ggml-metal-impl.h
+++ b/ggml/src/ggml-metal/ggml-metal-impl.h
@@ -120,6 +120,7 @@
#define FC_TOPK_MOE 1800
#define FC_MOE_REDUCE 1900
#define FC_DSV4_HC 2000
+#define FC_PAD 2100
// op-specific constants
#define OP_FLASH_ATTN_EXT_NQPSG 8
@@ -1120,6 +1121,10 @@ typedef struct {
uint64_t nb1;
uint64_t nb2;
uint64_t nb3;
+ int32_t lp0;
+ int32_t lp1;
+ int32_t lp2;
+ int32_t lp3;
} ggml_metal_kargs_pad;
typedef struct {
diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp
index 8a46ec66a..0ecd1a510 100644
--- a/ggml/src/ggml-metal/ggml-metal-ops.cpp
+++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp
@@ -4999,16 +4999,15 @@ int ggml_metal_op_pad(ggml_metal_op_t ctx, int idx) {
/*.nb0 =*/ nb0,
/*.nb1 =*/ nb1,
/*.nb2 =*/ nb2,
- /*.nb3 =*/ nb3
+ /*.nb3 =*/ nb3,
+ /*.lp0 =*/ ggml_get_op_params_i32(op, 0),
+ /*.lp1 =*/ ggml_get_op_params_i32(op, 2),
+ /*.lp2 =*/ ggml_get_op_params_i32(op, 4),
+ /*.lp3 =*/ ggml_get_op_params_i32(op, 6),
};
auto pipeline = ggml_metal_library_get_pipeline_pad(lib, op);
- if (pipeline.c4) {
- args.ne00 = ne00/4;
- args.ne0 = ne0/4;
- }
-
const int nth_max = MIN(64, ggml_metal_pipeline_max_theads_per_threadgroup(pipeline));
const int nth = MIN(args.ne0, nth_max);
const int nk0 = (args.ne0 + 1024 - 1)/1024; // note: 1024 is hardcoded in the kernel!
diff --git a/ggml/src/ggml-metal/kernels/misc.metal b/ggml/src/ggml-metal/kernels/misc.metal
index d3b01978f..d29786e1c 100644
--- a/ggml/src/ggml-metal/kernels/misc.metal
+++ b/ggml/src/ggml-metal/kernels/misc.metal
@@ -114,8 +114,14 @@ kernel void kernel_roll_f32(
}
}
-template <typename T>
-kernel void kernel_pad_impl(
+constant bool FC_pad_circular [[function_constant(FC_PAD + 0)]];
+
+// circular means on a torus, so the coordinates wrap around
+static inline int32_t wrap_around(int32_t coord, int32_t size) {
+ return (coord + size) % size;
+}
+
+kernel void kernel_pad_f32(
constant ggml_metal_kargs_pad & args,
device const char * src0,
device char * dst,
@@ -127,12 +133,40 @@ kernel void kernel_pad_impl(
const int32_t k0 = tgpig.x/args.ne1;
const int32_t i1 = tgpig.x - k0*args.ne1;
- const int32_t i03 = i3;
- const int32_t i02 = i2;
- const int32_t i01 = i1;
+ const int32_t ne00 = args.ne00;
+ const int32_t ne01 = args.ne01;
+ const int32_t ne02 = args.ne02;
+ const int32_t ne03 = args.ne03;
+
+ int32_t i01 = i1 - args.lp1;
+ int32_t i02 = i2 - args.lp2;
+ int32_t i03 = i3 - args.lp3;
- device const T * src0_ptr = (device const T *) (src0 + i03*args.nb03 + i02*args.nb02 + i01*args.nb01);
- device T * dst_ptr = (device T *) (dst + i3*args.nb3 + i2*args.nb2 + i1*args.nb1);
+ if (FC_pad_circular) {
+ i01 = wrap_around(i01, ne01);
+ i02 = wrap_around(i02, ne02);
+ i03 = wrap_around(i03, ne03);
+ }
+
+ device float * dst_ptr = (device float *) (dst + i3*args.nb3 + i2*args.nb2 + i1*args.nb1);
+
+ // the row lies in the padded region, so no source row backs it
+ if (i01 < 0 || i01 >= ne01 ||
+ i02 < 0 || i02 >= ne02 ||
+ i03 < 0 || i03 >= ne03) {
+ for (int32_t l0 = 0; l0 < 1024; l0 += ntg.x) {
+ const int32_t i0 = k0*1024 + tpitg.x + l0;
+ if (i0 >= args.ne0) {
+ break;
+ }
+
+ dst_ptr[i0] = 0.0f;
+ }
+
+ return;
+ }
+
+ device const char * src0_row = src0 + i03*args.nb03 + i02*args.nb02 + i01*args.nb01;
for (int32_t l0 = 0; l0 < 1024; l0 += ntg.x) {
const int32_t i0 = k0*1024 + tpitg.x + l0;
@@ -140,19 +174,16 @@ kernel void kernel_pad_impl(
break;
}
- if (i0 < args.ne00 && i1 < args.ne01 && i2 < args.ne02 && i3 < args.ne03) {
- dst_ptr[i0] = src0_ptr[i0];
- } else {
- dst_ptr[i0] = 0.0f;
+ int32_t i00 = i0 - args.lp0;
+
+ if (FC_pad_circular) {
+ i00 = wrap_around(i00, ne00);
}
+
+ dst_ptr[i0] = i00 >= 0 && i00 < ne00 ? *((device const float *) (src0_row + i00*args.nb00)) : 0.0f;
}
}
-typedef decltype(kernel_pad_impl<float>) kernel_pad_t;
-
-template [[host_name("kernel_pad_f32")]] kernel kernel_pad_t kernel_pad_impl<float>;
-template [[host_name("kernel_pad_f32_4")]] kernel kernel_pad_t kernel_pad_impl<float4>;
-
// TODO: this is slow - optimize
kernel void kernel_pad_reflect_1d_f32(
constant ggml_metal_kargs_pad_reflect_1d & args,
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index 510a29426..ff5a83295 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -10924,6 +10924,7 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
for (bool circular : {false, true}) {
test_cases.emplace_back(new test_pad_ext(GGML_TYPE_F32, {512, 512, 1, 1}, 0, 1, 0, 1, 0, 0, 0, 0, tfrm, circular));
test_cases.emplace_back(new test_pad_ext(GGML_TYPE_F32, {11, 22, 33, 44}, 1, 2, 3, 4, 5, 6, 7, 8, tfrm, circular));
+ test_cases.emplace_back(new test_pad_ext(GGML_TYPE_F32, {11, 22, 33, 44}, 0, 2, 0, 4, 0, 6, 0, 8, tfrm, circular));
}
}