Felipe97/llama-cpp-compiled
01.1k
1from __future__ import annotations2 3import json4 5from typing import Iterable, TYPE_CHECKING6 7import torch8 9if TYPE_CHECKING:10 from torch import Tensor11 12from .base import ModelBase, TextModel, gguf13 14 15@ModelBase.register("PlamoForCausalLM")16@ModelBase.example("pfnet/plamo-13b")17class PlamoModel(TextModel):18 model_arch = gguf.MODEL_ARCH.PLAMO19 20 def set_vocab(self):21 self._set_vocab_sentencepiece()22 23 def set_gguf_parameters(self):24 hparams = self.hparams25 26 self.gguf_writer.add_context_length(4096) # not in config.json27 self.gguf_writer.add_embedding_length(hparams["hidden_size"])28 self.gguf_writer.add_feed_forward_length(hparams["intermediate_size"])29 self.gguf_writer.add_block_count(self.block_count)30 self.gguf_writer.add_head_count(hparams["num_attention_heads"])31 self.gguf_writer.add_head_count_kv(5) # hparams["num_key_value_heads"]) is wrong32 self.gguf_writer.add_layer_norm_rms_eps(hparams["rms_norm_eps"])33 self.gguf_writer.add_file_type(self.ftype)34 35 def shuffle_attn_q_weight(self, data_torch):36 assert data_torch.size() == (5120, 5120)37 data_torch = data_torch.reshape(8, 5, 128, 5120)38 data_torch = torch.permute(data_torch, (1, 0, 2, 3))39 data_torch = torch.reshape(data_torch, (5120, 5120))40 return data_torch41 42 def shuffle_attn_output_weight(self, data_torch):43 assert data_torch.size() == (5120, 5120)44 data_torch = data_torch.reshape(5120, 8, 5, 128)45 data_torch = torch.permute(data_torch, (0, 2, 1, 3))46 data_torch = torch.reshape(data_torch, (5120, 5120))47 return data_torch48 49 def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:50 new_name = self.map_tensor_name(name)51 52 # shuffle for broadcasting of gqa in ggml_mul_mat53 if new_name.endswith("attn_q.weight"):54 data_torch = self.shuffle_attn_q_weight(data_torch)55 elif new_name.endswith("attn_output.weight"):56 data_torch = self.shuffle_attn_output_weight(data_torch)57 58 yield from super().modify_tensors(data_torch, name, bid)59 60 61@ModelBase.register("Plamo2ForCausalLM", "PLaMo2ForCausalLM")62@ModelBase.example("pfnet/plamo-2-1b")63class Plamo2Model(TextModel):64 model_arch = gguf.MODEL_ARCH.PLAMO265 66 def set_vocab(self):67 self._set_vocab_plamo()68 69 def set_gguf_parameters(self):70 hparams = self.hparams71 self.gguf_writer.add_vocab_size(self.hparams["vocab_size"])72 73 # Which layers are Mamba layers74 # PLaMo 2 uses mamba_step to indicate the pattern (e.g., 2 means every other layer)75 # This logic matches modeling_plamo.py's is_mamba function76 mamba_step = hparams.get("mamba_step", 2)77 mamba_enabled = hparams.get("mamba_enabled", True)78 num_key_value_heads = []79 num_attention_heads = []80 81 if mamba_enabled:82 for i in range(self.block_count):83 if self.block_count <= (mamba_step // 2):84 # use attention in last layer85 is_mamba = (i != self.block_count - 1)86 else:87 is_mamba = (i % mamba_step) != (mamba_step // 2)88 if is_mamba:89 num_key_value_heads.append(0)90 num_attention_heads.append(0)91 else:92 num_key_value_heads.append(hparams.get("num_key_value_heads", 4))93 num_attention_heads.append(hparams.get("num_attention_heads", 32))94 95 if num_key_value_heads and num_attention_heads:96 self.gguf_writer.add_head_count_kv(num_key_value_heads)97 self.gguf_writer.add_head_count(num_attention_heads)98 99 self.gguf_writer.add_context_length(hparams.get("max_position_embeddings", 2048))100 self.gguf_writer.add_embedding_length(hparams.get("hidden_size", 4096))101 self.gguf_writer.add_key_length(hparams.get("hidden_size_per_head", 128))102 self.gguf_writer.add_value_length(hparams.get("hidden_size_per_head", 128))103 self.gguf_writer.add_block_count(self.block_count)104 self.gguf_writer.add_layer_norm_rms_eps(hparams.get("rms_norm_eps", 1e-06))105 self.gguf_writer.add_rope_freq_base(self.rope_parameters.get("rope_theta", 10000))106 107 # Mamba parameters108 self.gguf_writer.add_ssm_state_size(hparams.get("mamba_d_state", 64))109 self.gguf_writer.add_ssm_conv_kernel(hparams.get("mamba_d_conv", 4))110 self.gguf_writer.add_ssm_time_step_rank(hparams.get("mamba_num_heads", 64))111 intermediate_size = hparams.get("mamba_num_heads", 64) * hparams.get("hidden_size_per_head", 128)112 self.gguf_writer.add_ssm_inner_size(intermediate_size)113 self.gguf_writer.add_ssm_group_count(0)114 115 # MLP feed forward parameters (for attention layers)116 self.gguf_writer.add_feed_forward_length(hparams.get("intermediate_size", 13312))117 self.gguf_writer.add_file_type(self.ftype)118 119 def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:120 if name.endswith(".A_log"):121 data_torch = -torch.exp(data_torch)122 elif name.endswith(".dt_bias"):123 name = name.rpartition(".dt_bias")[0] + ".dt_proj.bias"124 elif name.endswith(".dt_norm_weight"):125 name = name.rpartition(".dt_norm_weight")[0] + ".dt_norm.weight"126 elif name.endswith(".B_norm_weight"):127 name = name.rpartition(".B_norm_weight")[0] + ".B_norm.weight"128 elif name.endswith(".C_norm_weight"):129 name = name.rpartition(".C_norm_weight")[0] + ".C_norm.weight"130 elif name.endswith(".k_weight"):131 name = name.rpartition(".k_weight")[0] + ".k.weight"132 elif name.endswith(".q_weight"):133 name = name.rpartition(".q_weight")[0] + ".q.weight"134 elif name.endswith(".conv1d.weight"):135 data_torch = torch.squeeze(data_torch) # remove (, 1, )136 assert data_torch.ndim == 2137 elif name.endswith(".pre_mixer_norm.weight"):138 data_torch += 1.0139 elif name.endswith(".post_mixer_norm.weight"):140 data_torch += 1.0 / 5141 elif name.endswith(".pre_mlp_norm.weight"):142 data_torch += 1.0143 elif name.endswith(".post_mlp_norm.weight"):144 data_torch += 1.0 / (5**1.5)145 elif name.endswith(".norm.weight"):146 data_torch += 1.0147 148 yield from super().modify_tensors(data_torch, name, bid)149 150 151@ModelBase.register("Plamo3ForCausalLM", "PLaMo3ForCausalLM")152# [TAG_HF_EXAMPLE_GATED] pfnet/plamo-3-nict-2b-base is gated153@ModelBase.example("midorin-Linux/plamo-3-12b-self-merged-base")154class Plamo3Model(TextModel):155 model_arch = gguf.MODEL_ARCH.PLAMO3156 157 def set_vocab(self):158 self._set_vocab_plamo()159 160 tokenizer_config_path = self.dir_model / "tokenizer_config.json"161 tokenizer_config = {}162 163 if tokenizer_config_path.is_file():164 with open(tokenizer_config_path, encoding="utf-8") as f:165 tokenizer_config = json.load(f)166 167 chat_template = tokenizer_config.get("chat_template")168 chat_template_jinja = self.dir_model / "chat_template.jinja"169 170 if chat_template_jinja.is_file():171 with open(chat_template_jinja, encoding="utf-8") as f:172 chat_template = f.read()173 174 if chat_template:175 self.gguf_writer.add_chat_template(chat_template)176 177 def set_gguf_parameters(self):178 super().set_gguf_parameters()179 self.gguf_writer.add_vocab_size(self.hparams["vocab_size"])180 if (sliding_window := self.find_hparam(["window_size", "sliding_window"], optional=True)) is not None:181 self.gguf_writer.add_sliding_window(sliding_window)182 self.gguf_writer.add_sliding_window_pattern(self.hparams["sliding_window_pattern"])183 184 def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:185 186 if name.endswith(".pre_mixer_norm.weight"):187 data_torch = data_torch + 1.0188 elif name.endswith(".post_mixer_norm.weight"):189 data_torch = data_torch + 1.0 / 5190 elif name.endswith(".pre_mlp_norm.weight"):191 data_torch = data_torch + 1.0192 elif name.endswith(".post_mlp_norm.weight"):193 data_torch = data_torch + 1.0 / (5**1.5)194 elif name.endswith((".mixer.q_norm.weight", ".mixer.k_norm.weight")):195 data_torch = data_torch + 1.0196 elif name.endswith(".norm.weight"):197 data_torch = data_torch + 1.0198 199 yield from super().modify_tensors(data_torch, name, bid)200 