diff --git a/conversion/onyx.py b/conversion/onyx.py index e3acbc305e..01aae88b16 100644 --- a/conversion/onyx.py +++ b/conversion/onyx.py @@ -10,6 +10,18 @@ if TYPE_CHECKING: from .base import MmprojModel, ModelBase, TextModel, gguf +def _unpermute_for_rope(tensor: "Tensor", n_heads: int) -> "Tensor": + """Invert transformers' `_permute_for_rope`: HF stores Q/K in rotate_half layout, + llama.cpp consumes the interleaved (NORM) layout.""" + if tensor.ndim == 2: + dim1, dim2 = tensor.shape + return tensor.view(n_heads, 2, dim1 // n_heads // 2, dim2).transpose(1, 2).reshape(dim1, dim2) + if tensor.ndim == 1: + (dim1,) = tensor.shape + return tensor.view(n_heads, 2, dim1 // n_heads // 2).transpose(1, 2).reshape(dim1) + raise ValueError(f"_unpermute_for_rope: unexpected shape {tuple(tensor.shape)}") + + @ModelBase.register("OnyxForConditionalGeneration") class OnyxModel(TextModel): model_arch = gguf.MODEL_ARCH.ONYX @@ -46,6 +58,12 @@ class OnyxModel(TextModel): if shift != 0.0: data_torch = data_torch + shift + # Invert transformers' `_permute_for_rope` on Q/K, we keep ggml's NORM (interleaved) rope + if ".self_attn.q_proj." in name: + data_torch = _unpermute_for_rope(data_torch, int(self.hparams["num_attention_heads"])) + elif ".self_attn.k_proj." in name: + data_torch = _unpermute_for_rope(data_torch, int(self.hparams["num_key_value_heads"])) + # Synthesize QK-norm weights to absorb qk_scale_factor. # Onyx implementation: scaleless RMSNorm followed by qk_scale_factor.. if bid is not None and name.endswith(f"model.layers.{bid}.self_attn.q_proj.weight"): diff --git a/src/models/onyx.cpp b/src/models/onyx.cpp index ba2fceca0f..c0101a9891 100644 --- a/src/models/onyx.cpp +++ b/src/models/onyx.cpp @@ -66,7 +66,8 @@ llama_model_onyx::graph::graph(const llama_model & model, const llm_graph_params ggml_tensor * inpL; inpL = build_inp_embd(model.tok_embd); - // Onyx normalizes token embeddings with a scaleless RMSNorm that is already merged into the token embeddings in the transformers checkpoint. + inpL = build_norm(inpL, nullptr, nullptr, LLM_NORM_RMS, -1); + cb(inpL, "embd_norm", -1); ggml_tensor * inp_pos = build_inp_pos(); auto * inp_attn = build_attn_inp_kv_iswa();