Commit 9db0db6 for stable-diffusion.cpp
commit 9db0db65706fdf78544e52322943f8aed14f3d35
Author: Daniele <57776841+daniandtheweb@users.noreply.github.com>
Date: Tue Oct 6 10:54:55 2026 +0200
perf: avoid redundant permute+cont in non-interleaved RoPE (#2102)
diff --git a/src/model/common/rope.hpp b/src/model/common/rope.hpp
index ec82245..f0fd5b2 100644
--- a/src/model/common/rope.hpp
+++ b/src/model/common/rope.hpp
@@ -1082,8 +1082,16 @@ namespace Rope {
x = ggml_reshape_4d(ctx, x, 2, d_head / 2, L, n_head * N); // [N * n_head, L, d_head/2, 2]
x = ggml_cont(ctx, ggml_permute(ctx, x, 3, 0, 1, 2)); // [2, N * n_head, L, d_head/2]
} else {
- x = ggml_reshape_4d(ctx, x, d_head / 2, 2, L, n_head * N); // [N * n_head, L, 2, d_head/2]
- x = ggml_cont(ctx, ggml_ext_torch_permute(ctx, x, 0, 2, 3, 1)); // [2, N * n_head, L, d_head/2]
+ auto x_0 = ggml_view_3d(ctx, x, d_head / 2, L, n_head * N, x->nb[1], x->nb[2], 0);
+ auto x_1 = ggml_view_3d(ctx, x, d_head / 2, L, n_head * N, x->nb[1], x->nb[2], (d_head / 2) * x->nb[0]);
+
+ auto pe_2 = ggml_cont(ctx, ggml_permute(ctx, pe, 2, 3, 0, 1));
+ auto cos = ggml_view_2d(ctx, pe_2, d_head / 2, L, pe_2->nb[1], 0);
+ auto sin = ggml_view_2d(ctx, pe_2, d_head / 2, L, pe_2->nb[1], pe_2->nb[3]);
+
+ auto part_0 = ggml_sub(ctx, ggml_mul(ctx, x_0, cos), ggml_mul(ctx, x_1, sin));
+ auto part_1 = ggml_add(ctx, ggml_mul(ctx, x_0, sin), ggml_mul(ctx, x_1, cos));
+ return ggml_concat(ctx, part_0, part_1, 0);
}
int64_t offset = x->nb[2] * x->ne[2];
@@ -1101,10 +1109,7 @@ namespace Rope {
auto pe_1 = ggml_view_3d(ctx, pe, pe->ne[0], pe->ne[1], pe->ne[2], pe->nb[1], pe->nb[2], offset * 1); // [L, d_head/2, 2]
auto x_out = ggml_add_inplace(ctx, ggml_mul(ctx, x_0, pe_0), ggml_mul(ctx, x_1, pe_1)); // [N * n_head, L, d_head/2, 2]
- if (!rope_interleaved) {
- x_out = ggml_cont(ctx, ggml_permute(ctx, x_out, 1, 0, 2, 3)); // [N * n_head, L, x, d_head/2]
- }
- x_out = ggml_reshape_3d(ctx, x_out, d_head, L, n_head * N); // [N*n_head, L, d_head]
+ x_out = ggml_reshape_3d(ctx, x_out, d_head, L, n_head * N); // [N*n_head, L, d_head]
return x_out;
}