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("GPTNeoXForCausalLM")16@ModelBase.example("EleutherAI/pythia-70m")17class GPTNeoXModel(TextModel):18 model_arch = gguf.MODEL_ARCH.GPTNEOX19 20 def set_gguf_parameters(self):21 self.gguf_writer.add_context_length(self.hparams["max_position_embeddings"])22 self.gguf_writer.add_embedding_length(self.hparams["hidden_size"])23 self.gguf_writer.add_block_count(self.block_count)24 self.gguf_writer.add_feed_forward_length(self.hparams["intermediate_size"])25 self.gguf_writer.add_rope_dimension_count(26 int(self.hparams["rotary_pct"] * (self.hparams["hidden_size"] // self.hparams["num_attention_heads"])),27 )28 self.gguf_writer.add_head_count(self.hparams["num_attention_heads"])29 self.gguf_writer.add_parallel_residual(self.hparams.get("use_parallel_residual", True))30 self.gguf_writer.add_layer_norm_eps(self.hparams["layer_norm_eps"])31 32 def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:33 n_head = self.hparams.get("n_head", self.hparams.get("num_attention_heads"))34 n_embed = self.hparams.get("hidden_size", self.hparams.get("n_embed"))35 assert n_head is not None36 assert n_embed is not None37 38 if re.match(r"gpt_neox\.layers\.\d+\.attention\.query_key_value\.weight", name):39 # Map bloom-style qkv_linear to gpt-style qkv_linear40 # bloom: https://github.com/huggingface/transformers/blob/main/src/transformers/models/bloom/modeling_bloom.py#L238-L252 # noqa41 # gpt-2: https://github.com/huggingface/transformers/blob/main/src/transformers/models/gpt2/modeling_gpt2.py#L312 # noqa42 qkv_weights = data_torch.reshape((n_head, 3, n_embed // n_head, n_embed))43 data_torch = torch.cat(44 (45 qkv_weights[:, 0, :, :].reshape((-1, n_embed)),46 qkv_weights[:, 1, :, :].reshape((-1, n_embed)),47 qkv_weights[:, 2, :, :].reshape((-1, n_embed)),48 ),49 dim=0,50 )51 logger.info("re-format attention.linear_qkv.weight")52 elif re.match(r"gpt_neox\.layers\.\d+\.attention\.query_key_value\.bias", name):53 qkv_bias = data_torch.reshape((n_head, 3, n_embed // n_head))54 data_torch = torch.cat(55 (56 qkv_bias[:, 0, :].reshape((n_embed,)),57 qkv_bias[:, 1, :].reshape((n_embed,)),58 qkv_bias[:, 2, :].reshape((n_embed,)),59 ),60 dim=0,61 )62 logger.info("re-format attention.linear_qkv.bias")63 64 yield from super().modify_tensors(data_torch, name, bid)65 