This pattern appears in a lot of models, the rope operation is applied right before storing into the KV cache (usually on the K tensor). Add a path to some of the rope shaders that computes the destination address based on the set_rows tensor. Compile variants of the shader with D_TYPE of f16 (the usual KV cache type). Add a src3 operand to ggml_vk_op_f32 - sometimes rope uses three srcs and needs the fourth for the row indices. Add fused_ops_write_mask to indicate which intermediate tensors need to write their results to memory. Skipping writing the roped K value helps to allow more nodes to run concurrently. Add logic to ggml_vk_graph_optimize to make ROPE+VIEW+SET_ROWS consecutive. It rarely starts out that way in the graph. Add new backend tests.
49 lines
1.3 KiB
Plaintext
49 lines
1.3 KiB
Plaintext
#version 450
|
|
|
|
#include "rope_head.glsl"
|
|
|
|
void main() {
|
|
const uint i0 = 2*gl_GlobalInvocationID.y;
|
|
uint ne0 = p.ncols;
|
|
uint ne1 = p.p_delta_rows;
|
|
|
|
if (i0 >= ne0) {
|
|
return;
|
|
}
|
|
|
|
const uint row_dst = gl_GlobalInvocationID.x;
|
|
|
|
const uint row_x = row_dst % ne1;
|
|
const uint channel_x = row_dst / ne1;
|
|
|
|
uint idst = row_dst*ne0 + i0/2;
|
|
const uint ix = channel_x*p.s2 + row_x*p.s1 + i0/2;
|
|
|
|
// Fusion optimization: ROPE + VIEW + SET_ROWS..
|
|
// The rope output is viewed as a 1D tensor and offset based on a row index in data_i.
|
|
if (p.set_rows_stride != 0) {
|
|
idst = row_x*ne0 + i0/2;
|
|
idst += data_i[channel_x].x * p.set_rows_stride;
|
|
}
|
|
|
|
if (i0 >= p.n_dims) {
|
|
data_d[idst + i0/2 + 0] = D_TYPE(data_a[ix + i0/2 + 0]);
|
|
data_d[idst + i0/2 + 1] = D_TYPE(data_a[ix + i0/2 + 1]);
|
|
|
|
return;
|
|
}
|
|
|
|
const float theta_base = data_pos[channel_x] * pow(p.theta_scale, i0/2.0f);
|
|
|
|
const float freq_factor = p.has_ff != 0 ? data_ff[i0/2] : 1.0f;
|
|
|
|
float cos_theta, sin_theta;
|
|
rope_yarn(theta_base / freq_factor, i0, cos_theta, sin_theta);
|
|
|
|
const float x0 = float(data_a[ix + 0]);
|
|
const float x1 = float(data_a[ix + p.n_dims/2]);
|
|
|
|
data_d[idst + 0] = D_TYPE(x0*cos_theta - x1*sin_theta);
|
|
data_d[idst + p.n_dims/2] = D_TYPE(x0*sin_theta + x1*cos_theta);
|
|
}
|