diff --git a/src/model/common/rope.hpp b/src/model/common/rope.hpp index ec8224559..f0fd5b211 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; }