Felipe97/llama-cpp-compiled
01.1k
1from __future__ import annotations2 3import json4 5from typing import Iterable, TYPE_CHECKING6 7if TYPE_CHECKING:8 from torch import Tensor9 10from .base import MmprojModel, ModelBase, gguf, logger11 12from .llama import LlamaModel13 14 15@ModelBase.register(16 "LlavaForConditionalGeneration", # pixtral17 "Mistral3ForConditionalGeneration", # mistral small 3.118)19@ModelBase.example("mistral-community/pixtral-12b", "mistralai/Mistral-Small-3.1-24B-Instruct-2503")20class LlavaVisionModel(MmprojModel):21 img_break_tok_id = -122 use_break_tok = True23 24 def __init__(self, *args, **kwargs):25 super().__init__(*args, **kwargs)26 if self.hparams.get("model_type") == "pixtral":27 # layer_norm_eps is not in config.json, it is hard-coded in modeling_pixtral.py28 self.hparams["layer_norm_eps"] = self.hparams.get("layer_norm_eps", 1e-5)29 if self.use_break_tok:30 self.img_break_tok_id = self.get_token_id("[IMG_BREAK]")31 elif self.is_mistral_format:32 # hparams is already vision config here so norm_eps is only defined in global_config.33 self.hparams["norm_eps"] = self.global_config.get("norm_eps", None)34 assert self.hparams["norm_eps"] is not None, "norm_eps not found in params.json"35 if self.use_break_tok:36 self.img_break_tok_id = self.find_vparam(["image_break_token_id"])37 38 # params.json may ship -1 placeholders (Mistral Medium 3.5)39 # resolve the real id from the bundled tokenizer in that case40 if self.img_break_tok_id < 0:41 self.img_break_tok_id = self.get_mistral_token_id("[IMG_BREAK]")42 else:43 raise ValueError(f"Unsupported model type: {self.hparams['model_type']}")44 logger.info(f"Image break token id: {self.img_break_tok_id}")45 46 def get_token_id(self, token: str) -> int:47 tokenizer_config_file = self.dir_model / 'tokenizer_config.json'48 with open(tokenizer_config_file, "r", encoding="utf-8") as f:49 added_tokens_decoder = json.load(f).get('added_tokens_decoder') or {}50 for id_, token_data in added_tokens_decoder.items():51 if token_data.get("content") == token:52 return int(id_)53 # fallthrough to tokenizer.json54 with open(self.dir_model / "tokenizer.json", "r", encoding="utf-8") as f:55 tokenizer_json = json.load(f)56 for token_data in tokenizer_json["added_tokens"]:57 if token_data["content"] == token:58 return int(token_data["id"])59 raise ValueError(f"Token '{token}' not found in tokenizer config.")60 61 def get_mistral_token_id(self, token: str) -> int:62 # mistral native format ships tekken.json or a versioned spm tokenizer63 tekken_file = self.dir_model / "tekken.json"64 if tekken_file.is_file():65 with open(tekken_file, "r", encoding="utf-8") as f:66 data = json.load(f)67 for entry in data.get("special_tokens", []):68 if entry.get("token_str") == token:69 return int(entry["rank"])70 tokenizer_json_file = self.dir_model / "tokenizer.json"71 if tokenizer_json_file.is_file():72 with open(tokenizer_json_file, "r", encoding="utf-8") as f:73 data = json.load(f)74 for entry in data.get("added_tokens", []):75 if entry.get("content") == token:76 return int(entry["id"])77 raise ValueError(f"Token '{token}' not found in mistral tokenizer files.")78 79 def set_gguf_parameters(self):80 super().set_gguf_parameters()81 hparams = self.hparams82 if hparams.get("model_type") == "pixtral":83 self.gguf_writer.add_clip_projector_type(gguf.VisionProjectorType.PIXTRAL)84 self.gguf_writer.add_vision_attention_layernorm_eps(hparams["layer_norm_eps"])85 86 # hidden_act87 if hparams["hidden_act"] == "silu":88 self.gguf_writer.add_vision_use_silu(True)89 elif hparams["hidden_act"] == "gelu":90 self.gguf_writer.add_vision_use_gelu(True)91 else:92 raise ValueError(f"Unsupported hidden_act: {hparams['hidden_act']}")93 94 # spatial_merge_size95 if "spatial_merge_size" in self.global_config:96 self.gguf_writer.add_vision_spatial_merge_size(self.global_config["spatial_merge_size"])97 98 def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:99 n_head = (100 self.hparams["num_attention_heads"] if not self.is_mistral_format else self.find_vparam(["num_attention_heads"])101 )102 n_kv_head = n_head103 104 valid_prefixes = (105 "multi_modal_projector.",106 "vision_tower.",107 "vision_encoder.",108 "vision_language_adapter.",109 "patch_merger.",110 "pre_mm_projector_norm",111 )112 113 if any(name.startswith(prefix) for prefix in valid_prefixes):114 # process vision tensors115 if name.endswith(("q_proj.weight", "q_proj.bias")) and not self.is_mistral_format:116 data_torch = LlamaModel.permute(data_torch, n_head, n_head)117 if name.endswith(("k_proj.weight", "k_proj.bias")) and not self.is_mistral_format:118 data_torch = LlamaModel.permute(data_torch, n_head, n_kv_head)119 yield from super().modify_tensors(data_torch, name, bid)120 return121 122 embed_key = "embed_tokens.weight" if not self.is_mistral_format else "tok_embeddings.weight"123 if self.img_break_tok_id > 0 and embed_key in name:124 logger.info(f"Extracting [IMG_BREAK] token embedding from {name}")125 # for pixtral model, we need to extract the [IMG_BREAK] token embedding126 img_break_embd = data_torch[self.img_break_tok_id]127 name = gguf.TENSOR_NAMES[gguf.MODEL_TENSOR.V_TOK_EMBD_IMG_BREAK]128 yield from super().modify_tensors(img_break_embd, name, bid)129 130 return # skip other tensors131 