Felipe97/llama-cpp-compiled
01.2k
1from __future__ import annotations2 3from typing import Callable, Iterable, TYPE_CHECKING4 5if TYPE_CHECKING:6 from torch import Tensor7 8from .base import ModelBase, TextModel, gguf9 10from .llama import LlamaModel11 12 13@ModelBase.register("ChameleonForConditionalGeneration")14@ModelBase.register("ChameleonForCausalLM") # obsolete15# [TAG_HF_EXAMPLE_GATED] facebook/chameleon-7b is gated16# [TAG_HF_EXAMPLE_MISSING]17class ChameleonModel(TextModel):18 model_arch = gguf.MODEL_ARCH.CHAMELEON19 20 def set_gguf_parameters(self):21 super().set_gguf_parameters()22 self.gguf_writer.add_swin_norm(self.hparams.get("swin_norm", False))23 24 def set_vocab(self):25 self._set_vocab_gpt2()26 27 @classmethod28 def filter_tensors(cls, item: tuple[str, Callable[[], Tensor]]) -> tuple[str, Callable[[], Tensor]] | None:29 name, gen = item30 31 # ignore image tokenizer for now32 # TODO: image support for Chameleon33 if name.startswith("model.vqmodel"):34 return None35 36 return super().filter_tensors(item)37 38 def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:39 n_head = self.hparams["num_attention_heads"]40 n_kv_head = self.hparams.get("num_key_value_heads")41 hidden_dim = self.hparams.get("hidden_size")42 43 if name.endswith(("q_proj.weight", "q_proj.bias")):44 data_torch = LlamaModel.permute(data_torch, n_head, n_head)45 if name.endswith(("k_proj.weight", "k_proj.bias")):46 data_torch = LlamaModel.permute(data_torch, n_head, n_kv_head)47 if name.endswith(("q_norm.weight", "q_norm.bias")):48 data_torch = ChameleonModel._reverse_hf_permute(data_torch, n_head, hidden_dim)49 if name.endswith(("k_norm.weight", "k_norm.bias")):50 data_torch = ChameleonModel._reverse_hf_permute(data_torch, n_kv_head, hidden_dim)51 52 yield from super().modify_tensors(data_torch, name, bid)53 54 # see: https://github.com/huggingface/transformers/blob/72fb02c47dbbe1999ae105319f24631cad6e2e00/src/transformers/models/chameleon/convert_chameleon_weights_to_hf.py#L176-L20355 @staticmethod56 def _reverse_hf_permute(data_torch, n_heads, hidden_dim):57 head_dim = hidden_dim // n_heads58 data_torch = data_torch[0].view(2, head_dim // 2).t().reshape(1, -1)59 data_torch = data_torch.repeat_interleave(n_heads, 0)60 return data_torch61 