From e56d4d4b7ca7ded3d3c336a427fe4a596066eea6 Mon Sep 17 00:00:00 2001 From: Daniele <57776841+daniandtheweb@users.noreply.github.com> Date: Sun, 4 Oct 2026 13:14:55 +0200 Subject: [PATCH] perf: avoid redundant permute+cont in non-interleaved RoPE --- src/model/common/rope.hpp | 60 +++++++++++++++++++++++---------------- 1 file changed, 36 insertions(+), 24 deletions(-) diff --git a/src/model/common/rope.hpp b/src/model/common/rope.hpp index ec8224559..8c74e639a 100644 --- a/src/model/common/rope.hpp +++ b/src/model/common/rope.hpp @@ -1073,39 +1073,51 @@ namespace Rope { bool rope_interleaved = true) { // x: [N, L, n_head, d_head] // pe: [L, d_head/2, 2, 2], [[cos, -sin], [sin, cos]] + // + // Returns [N * n_head, L, d_head]. int64_t d_head = x->ne[0]; int64_t n_head = x->ne[1]; int64_t L = x->ne[2]; int64_t N = x->ne[3]; - x = ggml_cont(ctx, ggml_permute(ctx, x, 0, 2, 1, 3)); // [N, n_head, L, d_head] + + x = ggml_cont(ctx, ggml_permute(ctx, x, 0, 2, 1, 3)); // [N, n_head, L, d_head] + if (rope_interleaved) { 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] - } - int64_t offset = x->nb[2] * x->ne[2]; - auto x_0 = ggml_view_3d(ctx, x, x->ne[0], x->ne[1], x->ne[2], x->nb[1], x->nb[2], offset * 0); // [N * n_head, L, d_head/2] - auto x_1 = ggml_view_3d(ctx, x, x->ne[0], x->ne[1], x->ne[2], x->nb[1], x->nb[2], offset * 1); // [N * n_head, L, d_head/2] - x_0 = ggml_reshape_4d(ctx, x_0, 1, x_0->ne[0], x_0->ne[1], x_0->ne[2]); // [N * n_head, L, d_head/2, 1] - x_1 = ggml_reshape_4d(ctx, x_1, 1, x_1->ne[0], x_1->ne[1], x_1->ne[2]); // [N * n_head, L, d_head/2, 1] - auto temp_x = ggml_new_tensor_4d(ctx, x_0->type, 2, x_0->ne[1], x_0->ne[2], x_0->ne[3]); - x_0 = ggml_repeat(ctx, x_0, temp_x); // [N * n_head, L, d_head/2, 2] - x_1 = ggml_repeat(ctx, x_1, temp_x); // [N * n_head, L, d_head/2, 2] - - pe = ggml_cont(ctx, ggml_permute(ctx, pe, 3, 0, 1, 2)); // [2, L, d_head/2, 2] - offset = pe->nb[2] * pe->ne[2]; - auto pe_0 = ggml_view_3d(ctx, pe, pe->ne[0], pe->ne[1], pe->ne[2], pe->nb[1], pe->nb[2], offset * 0); // [L, d_head/2, 2] - 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] + int64_t offset = x->nb[2] * x->ne[2]; + auto x_0 = ggml_view_3d(ctx, x, x->ne[0], x->ne[1], x->ne[2], x->nb[1], x->nb[2], offset * 0); // [N * n_head, L, d_head/2] + auto x_1 = ggml_view_3d(ctx, x, x->ne[0], x->ne[1], x->ne[2], x->nb[1], x->nb[2], offset * 1); // [N * n_head, L, d_head/2] + x_0 = ggml_reshape_4d(ctx, x_0, 1, x_0->ne[0], x_0->ne[1], x_0->ne[2]); // [N * n_head, L, d_head/2, 1] + x_1 = ggml_reshape_4d(ctx, x_1, 1, x_1->ne[0], x_1->ne[1], x_1->ne[2]); // [N * n_head, L, d_head/2, 1] + auto temp_x = ggml_new_tensor_4d(ctx, x_0->type, 2, x_0->ne[1], x_0->ne[2], x_0->ne[3]); + x_0 = ggml_repeat(ctx, x_0, temp_x); // [N * n_head, L, d_head/2, 2] + x_1 = ggml_repeat(ctx, x_1, temp_x); // [N * n_head, L, d_head/2, 2] + + pe = ggml_cont(ctx, ggml_permute(ctx, pe, 3, 0, 1, 2)); // [2, L, d_head/2, 2] + offset = pe->nb[2] * pe->ne[2]; + auto pe_0 = ggml_view_3d(ctx, pe, pe->ne[0], pe->ne[1], pe->ne[2], pe->nb[1], pe->nb[2], offset * 0); // [L, d_head/2, 2] + 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] + x_out = ggml_reshape_3d(ctx, x_out, d_head, L, n_head * N); // [N*n_head, L, d_head] + return x_out; } - x_out = ggml_reshape_3d(ctx, x_out, d_head, L, n_head * N); // [N*n_head, L, d_head] - return x_out; + + // Non-interleaved: pairs (i, i + d_head/2) are contiguous halves of the head dim, so the + // rotation is computed on views of the permuted x and assembled with one concat, avoiding + // the extra permute+cont round trips the interleaved layout requires. + auto x_lo = ggml_view_3d(ctx, x, d_head / 2, L, n_head * N, x->nb[1], x->nb[2], 0); // [d_head/2, L, n_head*N] + auto x_hi = ggml_view_3d(ctx, x, d_head / 2, L, n_head * N, x->nb[1], x->nb[2], (d_head / 2) * x->nb[0]); // [d_head/2, L, n_head*N] + + auto pe2 = ggml_cont(ctx, ggml_permute(ctx, pe, 2, 3, 0, 1)); // [d_head/2, L, 2, 2] + auto cos = ggml_view_2d(ctx, pe2, d_head / 2, L, pe2->nb[1], 0); // [d_head/2, L] = pe[0,0,i,l] + auto sin = ggml_view_2d(ctx, pe2, d_head / 2, L, pe2->nb[1], pe2->nb[3]); // [d_head/2, L] = pe[0,1,i,l] + + auto part0 = ggml_sub(ctx, ggml_mul(ctx, x_lo, cos), ggml_mul(ctx, x_hi, sin)); // [d_head/2, L, n_head*N] + auto part1 = ggml_add(ctx, ggml_mul(ctx, x_lo, sin), ggml_mul(ctx, x_hi, cos)); // [d_head/2, L, n_head*N] + return ggml_concat(ctx, part0, part1, 0); // [d_head, L, n_head*N] } __STATIC_INLINE__ ggml_tensor* attention(GGMLRunnerContext* ctx,