Felipe97/llama-cpp-compiled
01.1k
1from __future__ import annotations2 3import re4 5from typing import Iterable, TYPE_CHECKING6 7import torch8 9if TYPE_CHECKING:10 from torch import Tensor11 12from .base import ModelBase, TextModel, gguf, logger13 14 15@ModelBase.register("BloomForCausalLM", "BloomModel")16@ModelBase.example("bigscience/bloom-560m")17class BloomModel(TextModel):18 model_arch = gguf.MODEL_ARCH.BLOOM19 20 def set_gguf_parameters(self):21 n_embed = self.hparams.get("hidden_size", self.hparams.get("n_embed"))22 n_head = self.hparams.get("n_head", self.hparams.get("num_attention_heads"))23 assert n_head is not None24 assert n_embed is not None25 self.gguf_writer.add_context_length(self.hparams.get("seq_length", n_embed))26 self.gguf_writer.add_embedding_length(n_embed)27 self.gguf_writer.add_feed_forward_length(4 * n_embed)28 self.gguf_writer.add_block_count(self.block_count)29 self.gguf_writer.add_head_count(n_head)30 self.gguf_writer.add_head_count_kv(n_head)31 self.gguf_writer.add_layer_norm_eps(self.hparams["layer_norm_epsilon"])32 self.gguf_writer.add_file_type(self.ftype)33 34 def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:35 n_head = self.hparams.get("n_head", self.hparams.get("num_attention_heads"))36 n_embed = self.hparams.get("hidden_size", self.hparams.get("n_embed"))37 assert n_head is not None38 assert n_embed is not None39 40 name = re.sub(r'transformer\.', '', name)41 42 if re.match(r"h\.\d+\.self_attention\.query_key_value\.weight", name):43 # Map bloom-style qkv_linear to gpt-style qkv_linear44 # bloom: https://github.com/huggingface/transformers/blob/main/src/transformers/models/bloom/modeling_bloom.py#L238-L252 # noqa45 # gpt-2: https://github.com/huggingface/transformers/blob/main/src/transformers/models/gpt2/modeling_gpt2.py#L312 # noqa46 qkv_weights = data_torch.reshape((n_head, 3, n_embed // n_head, n_embed))47 data_torch = torch.cat(48 (49 qkv_weights[:, 0, :, :].reshape((-1, n_embed)),50 qkv_weights[:, 1, :, :].reshape((-1, n_embed)),51 qkv_weights[:, 2, :, :].reshape((-1, n_embed)),52 ),53 dim=0,54 )55 logger.info("re-format attention.linear_qkv.weight")56 elif re.match(r"h\.\d+\.self_attention\.query_key_value\.bias", name):57 qkv_bias = data_torch.reshape((n_head, 3, n_embed // n_head))58 data_torch = torch.cat(59 (60 qkv_bias[:, 0, :].reshape((n_embed,)),61 qkv_bias[:, 1, :].reshape((n_embed,)),62 qkv_bias[:, 2, :].reshape((n_embed,)),63 ),64 dim=0,65 )66 logger.info("re-format attention.linear_qkv.bias")67 68 yield from super().modify_tensors(data_torch, name, bid)69 