n_patches(n_patches_x * n_patches_y),
n_embd(hparams.n_embd),
n_head(hparams.n_head),
+ n_head_kv(hparams.n_head_kv),
d_head(n_embd / n_head),
n_layer(hparams.n_layer),
n_mmproj_embd(clip_n_mmproj_embd(ctx)),
}
}
- Qcur = ggml_reshape_3d(ctx0, Qcur, d_head, n_head, n_pos);
- Kcur = ggml_reshape_3d(ctx0, Kcur, d_head, n_head, n_pos);
- Vcur = ggml_reshape_3d(ctx0, Vcur, d_head, n_head, n_pos);
+ Qcur = ggml_reshape_3d(ctx0, Qcur, d_head, n_head, n_pos);
+ Kcur = ggml_reshape_3d(ctx0, Kcur, d_head, n_head_kv, n_pos);
+ Vcur = ggml_reshape_3d(ctx0, Vcur, d_head, n_head_kv, n_pos);
if (norm_per_head) {
if (layer.q_norm) {
get_u32(string_format(KEY_PROJ_DIM, prefix), hparams.projection_dim);
get_f32(string_format(KEY_LAYER_NORM_EPS, prefix), hparams.eps);
+ // n_head_kv is optional (for GQA), default to n_head
+ hparams.n_head_kv = hparams.n_head;
+
if (is_vision) {
get_u32(KEY_IMAGE_SIZE, hparams.image_size);
get_u32(KEY_PATCH_SIZE, hparams.patch_size);