CoolFace
Modelpublic

Felipe97/llama-cpp-compiled

sourceHugging Faceupdated 3d agoView on Hugging Face
0likes1.1kdownloads
stablelm.py100 linesDownload Raw Back to conversion
1from __future__ import annotations2 3from typing import Iterable, TYPE_CHECKING4 5import torch6 7if TYPE_CHECKING:8    from torch import Tensor9 10from .base import ModelBase, TextModel, gguf11 12 13@ModelBase.register("StableLmForCausalLM", "StableLMEpochForCausalLM", "LlavaStableLMEpochForCausalLM")14@ModelBase.example("stabilityai/stablelm-2-1_6b")15class StableLMModel(TextModel):16    model_arch = gguf.MODEL_ARCH.STABLELM17 18    def set_vocab(self):19        if (self.dir_model / "tokenizer.json").is_file():20            self._set_vocab_gpt2()21        else:22            # StableLM 2 1.6B used to have a vocab in a similar format to Qwen's vocab23            self._set_vocab_qwen()24 25    def set_gguf_parameters(self):26        hparams = self.hparams27 28        self.gguf_writer.add_context_length(hparams["max_position_embeddings"])29        self.gguf_writer.add_embedding_length(hparams["hidden_size"])30        self.gguf_writer.add_block_count(self.block_count)31        self.gguf_writer.add_feed_forward_length(hparams["intermediate_size"])32        rotary_factor = self.rope_parameters["partial_rotary_factor"]33        self.gguf_writer.add_rope_dimension_count(int(rotary_factor * (hparams["hidden_size"] // hparams["num_attention_heads"])))34        self.gguf_writer.add_head_count(hparams["num_attention_heads"])35        self.gguf_writer.add_head_count_kv(hparams["num_key_value_heads"])36        self.gguf_writer.add_parallel_residual(hparams["use_parallel_residual"] if "use_parallel_residual" in hparams else True)37        self.gguf_writer.add_layer_norm_eps(self.find_hparam(["layer_norm_eps", "norm_eps"]))38        self.gguf_writer.add_file_type(self.ftype)39 40    _q_norms: list[dict[str, Tensor]] | None = None41    _k_norms: list[dict[str, Tensor]] | None = None42 43    def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:44        n_head = self.hparams["num_attention_heads"]45        n_kv_head = self.hparams["num_key_value_heads"]46 47        if name.find("q_layernorm.norms") != -1:48            assert bid is not None49 50            if self._q_norms is None:51                self._q_norms = [{} for _ in range(self.block_count)]52 53            self._q_norms[bid][name] = data_torch54 55            if len(self._q_norms[bid]) >= n_head:56                return self._stack_qk_norm(bid, n_head, self._q_norms[bid], "q_layernorm")57            else:58                return59 60        if name.find("k_layernorm.norms") != -1:61            assert bid is not None62 63            if self._k_norms is None:64                self._k_norms = [{} for _ in range(self.block_count)]65 66            self._k_norms[bid][name] = data_torch67 68            if len(self._k_norms[bid]) >= n_kv_head:69                return self._stack_qk_norm(bid, n_kv_head, self._k_norms[bid], "k_layernorm")70            else:71                return72 73        yield from super().modify_tensors(data_torch, name, bid)74 75    def _stack_qk_norm(self, bid: int, n_head: int, norms: dict[str, Tensor], layer_name: str = "q_layernorm"):76        datas: list[Tensor] = []77        # extract the norms in order78        for xid in range(n_head):79            ename = f"model.layers.{bid}.self_attn.{layer_name}.norms.{xid}.weight"80            datas.append(norms[ename])81            del norms[ename]82        data_torch = torch.stack(datas, dim=0)83 84        merged_name = f"model.layers.{bid}.self_attn.{layer_name}.weight"85 86        yield from super().modify_tensors(data_torch, merged_name, bid)87 88    def prepare_tensors(self):89        super().prepare_tensors()90 91        if self._q_norms is not None or self._k_norms is not None:92            # flatten two `list[dict[str, Tensor]]` into a single `list[str]`93            norms = (94                [k for d in self._q_norms for k in d.keys()] if self._q_norms is not None else []95            ) + (96                [k for d in self._k_norms for k in d.keys()] if self._k_norms is not None else []97            )98            if len(norms) > 0:99                raise ValueError(f"Unprocessed norms: {norms}")100