CoolFace
Modelpublic

Felipe97/llama-cpp-compiled

sourceHugging Faceupdated 3d agoView on Hugging Face
0likes1.1kdownloads
mamba.py201 linesDownload Raw Back to conversion
1from __future__ import annotations2 3import json4 5from pathlib import Path6from typing import Callable, Iterable, TYPE_CHECKING7 8import torch9 10if TYPE_CHECKING:11    from torch import Tensor12 13from .base import ModelBase, TextModel, gguf, logger14 15 16@ModelBase.register("MambaForCausalLM", "MambaLMHeadModel", "FalconMambaForCausalLM")17@ModelBase.example("state-spaces/mamba-130m-hf", "tiiuae/falcon-mamba-7b")18class MambaModel(TextModel):19    model_arch = gguf.MODEL_ARCH.MAMBA20 21    def __init__(self, dir_model: Path, *args, **kwargs):22        # Avoid using AutoConfig for hparams23        hparams = kwargs.pop("hparams", None)24        if hparams is None:25            with open(dir_model / "config.json", "r", encoding="utf-8") as f:26                hparams = json.load(f)27        super().__init__(dir_model, *args, hparams=hparams, **kwargs)28 29    def set_vocab(self):30        vocab_size = self.hparams["vocab_size"]31        # Round vocab size to next multiple of 832        pad_vocab = self.hparams.get("pad_vocab_size_multiple", 8)33        # pad using ceiling division34        # ref: https://stackoverflow.com/a/17511341/2282786335        vocab_size = -(vocab_size // -pad_vocab) * pad_vocab36        self.hparams["vocab_size"] = vocab_size37 38        if (self.dir_model / "tokenizer.json").is_file():39            self._set_vocab_gpt2()40        elif (self.dir_model / "tokenizer.model").is_file():41            self._set_vocab_sentencepiece()42        else:43            # Use the GPT-NeoX tokenizer when no tokenizer files are present44            self._set_vocab_builtin("gpt-neox", vocab_size)45 46    def set_gguf_parameters(self):47        d_model = self.find_hparam(["hidden_size",       "d_model"])48        d_conv  = self.find_hparam(["conv_kernel",       "d_conv"],  optional=True) or 449        d_inner = self.find_hparam(["intermediate_size", "d_inner"], optional=True) or 2 * d_model50        d_state = self.find_hparam(["state_size",        "d_state"], optional=True) or 1651        # ceiling division52        # ref: https://stackoverflow.com/a/17511341/2282786353        # ref: https://github.com/state-spaces/mamba/blob/ce59daea3a090d011d6476c6e5b97f6d58ddad8b/mamba_ssm/modules/mamba_simple.py#L5854        dt_rank      = self.find_hparam(["time_step_rank",     "dt_rank"],      optional=True) or -(d_model // -16)55        rms_norm_eps = self.find_hparam(["layer_norm_epsilon", "rms_norm_eps"], optional=True) or 1e-556        use_dt_b_c_norm = False57        # For falconmamba we do apply RMS norm on B / DT and C layers58        if self.find_hparam(["model_type"], optional=True) in ("falcon_mamba",):59            use_dt_b_c_norm = True60        # Fail early for models which don't have a block expansion factor of 261        assert d_inner == 2 * d_model62 63        self.gguf_writer.add_context_length(2**20) # arbitrary value; for those who use the default64        self.gguf_writer.add_embedding_length(d_model)65        self.gguf_writer.add_feed_forward_length(0) # unused, but seemingly required when loading66        self.gguf_writer.add_head_count(0) # unused, but seemingly required when loading67        self.gguf_writer.add_block_count(self.block_count)68        self.gguf_writer.add_ssm_conv_kernel(d_conv)69        self.gguf_writer.add_ssm_inner_size(d_inner)70        self.gguf_writer.add_ssm_state_size(d_state)71        self.gguf_writer.add_ssm_time_step_rank(dt_rank)72        self.gguf_writer.add_layer_norm_rms_eps(rms_norm_eps)73        self.gguf_writer.add_ssm_dt_b_c_rms(use_dt_b_c_norm) # For classic Mamba we don't apply rms norm on B / DT layers74        self.gguf_writer.add_file_type(self.ftype)75 76    _tok_embd = None77 78    def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:79        output_name = self.format_tensor_name(gguf.MODEL_TENSOR.OUTPUT)80        tok_embd_name = self.format_tensor_name(gguf.MODEL_TENSOR.TOKEN_EMBD)81 82        new_name = self.map_tensor_name(name)83 84        if name.endswith(".A_log"):85            logger.debug("A_log --> A ==> " + new_name)86            data_torch = -torch.exp(data_torch)87 88        # [4 1 8192 1] -> [4 8192 1 1]89        if self.match_model_tensor_name(new_name, gguf.MODEL_TENSOR.SSM_CONV1D, bid):90            data_torch = data_torch.squeeze()91 92        # assuming token_embd.weight is seen before output.weight93        if self._tok_embd is not None and new_name == output_name:94            if torch.equal(self._tok_embd, data_torch):95                logger.debug(f"{output_name} is equivalent to {tok_embd_name}, omitting")96                return97        elif new_name == tok_embd_name:98            self._tok_embd = data_torch99 100        yield from super().modify_tensors(data_torch, new_name, bid)101 102 103@ModelBase.register("Mamba2ForCausalLM")104@ModelBase.example("mistralai/Mamba-Codestral-7B-v0.1")105class Mamba2Model(TextModel):106    model_arch = gguf.MODEL_ARCH.MAMBA2107 108    def __init__(self, dir_model: Path, *args, **kwargs):109        # Avoid using AutoConfig for hparams110        # It wrongly assumes all Mamba2 models are Mamba-Codestral-7B-v0.1111        hparams = kwargs.pop("hparams", None)112        if hparams is None:113            with open(dir_model / "config.json", "r", encoding="utf-8") as f:114                hparams = json.load(f)115        if "llm_config" in hparams:116            hparams["text_config"] = hparams["llm_config"]117        super().__init__(dir_model, *args, hparams=hparams, **kwargs)118        self.d_model = self.find_hparam(["hidden_size", "d_model", "dim"])119        self.expand = self.find_hparam(["mamba_expand", "expand"], optional=True) or 2120        self.d_inner = self.find_hparam(["mamba_d_ssm", "intermediate_size", "d_inner"], optional=True) or self.expand * self.d_model121        self.n_group = self.find_hparam(["n_groups"], optional=True) or 1122 123    def set_vocab(self):124        vocab_size = self.hparams["vocab_size"]125        # Round vocab size to next multiple of 16126        pad_vocab = self.hparams.get("pad_vocab_size_multiple", 16)127        # pad using ceiling division128        # ref: https://stackoverflow.com/a/17511341/22827863129        vocab_size = -(vocab_size // -pad_vocab) * pad_vocab130        self.hparams["vocab_size"] = vocab_size131 132        if (self.dir_model / "tokenizer.model").is_file():133            self._set_vocab_sentencepiece()134        elif (self.dir_model / "tokenizer.model.v3").is_file():135            # mamba-codestral136            raise NotImplementedError(f"Please rename {self.dir_model / 'tokenizer.model.v3'} to {self.dir_model / 'tokenizer.model'}")137        elif (self.dir_model / "tokenizer.json").is_file():138            self._set_vocab_gpt2()139        else:140            # Use the GPT-NeoX tokenizer when no tokenizer files are present141            self._set_vocab_builtin("gpt-neox", vocab_size)142 143    def set_gguf_parameters(self):144        d_conv  = self.find_hparam(["conv_kernel", "d_conv"],     optional=True) or 4145        d_state = self.find_hparam(["state_size",  "d_state"],    optional=True) or 128146        head_dim = self.find_hparam(["mamba_d_head", "head_dim"], optional=True) or 64147 148        rms_norm_eps = self.find_hparam(["layer_norm_epsilon", "rms_norm_eps"], optional=True) or 1e-5149 150        # skip the assertion for FalconH1 Model151        if self.model_arch != gguf.MODEL_ARCH.FALCON_H1:152            assert self.d_inner == self.expand * self.d_model153            assert self.d_inner % head_dim == 0154 155        self.gguf_writer.add_context_length(2**20)  # arbitrary value; for those who use the default156        self.gguf_writer.add_embedding_length(self.d_model)157        self.gguf_writer.add_feed_forward_length(0)  # unused, but seemingly required when loading158        self.gguf_writer.add_head_count(0)  # unused, but seemingly required when loading159        self.gguf_writer.add_block_count(self.block_count)160        self.gguf_writer.add_ssm_conv_kernel(d_conv)161        self.gguf_writer.add_ssm_inner_size(self.d_inner)162        self.gguf_writer.add_ssm_state_size(d_state)163        self.gguf_writer.add_ssm_time_step_rank(self.d_inner // head_dim)164        self.gguf_writer.add_ssm_group_count(self.n_group)165        self.gguf_writer.add_layer_norm_rms_eps(rms_norm_eps)166        self.gguf_writer.add_file_type(self.ftype)167 168    @classmethod169    def filter_tensors(cls, item: tuple[str, Callable[[], Tensor]]) -> tuple[str, Callable[[], Tensor]] | None:170        name, gen = item171 172        if name.startswith(("model.backbone", "model.lm_head")):173            # map Mamba-Codestral-7B-v0.1 tensor names to the names used by Mamba-2174            name = name.removeprefix("model.")175 176        if name.endswith(".dt_bias"):177            name = name.rpartition(".dt_bias")[0] + ".dt_proj.bias"178 179        return super().filter_tensors((name, gen))180 181    def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:182        new_name = self.map_tensor_name(name)183 184        if self.match_model_tensor_name(new_name, gguf.MODEL_TENSOR.SSM_CONV1D, bid):185            data_torch = data_torch.squeeze()186        elif any(self.match_model_tensor_name(new_name, t, bid, suffix="") for t in [187            gguf.MODEL_TENSOR.SSM_A,188            gguf.MODEL_TENSOR.SSM_D,189        ]):190            # unsqueeze A to use similar shape semantics as Mamba-1191            # (D is also unsqueezed, but for more straightforward broadcast internally)192            data_torch = data_torch.reshape((*data_torch.shape, 1))193        elif self.match_model_tensor_name(new_name, gguf.MODEL_TENSOR.SSM_NORM, bid):194            data_torch = data_torch.reshape((self.n_group, self.d_inner // self.n_group))195 196        if name.endswith(".A_log"):197            logger.debug("A_log --> A ==> " + new_name)198            data_torch = -torch.exp(data_torch)199 200        yield (new_name, data_torch)201