Felipe97/llama-cpp-compiled
01.1k
1from __future__ import annotations2 3from typing import Callable, Iterable, TYPE_CHECKING4 5import torch6 7if TYPE_CHECKING:8 from torch import Tensor9 10from .base import ModelBase, TextModel, gguf, logger11 12 13@ModelBase.register("GptOssForCausalLM")14@ModelBase.example("openai/gpt-oss-20b")15class GptOssModel(TextModel):16 model_arch = gguf.MODEL_ARCH.GPT_OSS17 18 # TODO: remove once MXFP4 is supported more generally19 def dequant_model(self):20 if self._is_mxfp4:21 return22 return super().dequant_model()23 24 def transform_nibble_layout(self, tensor):25 assert tensor.dtype == torch.uint826 assert tensor.shape[-1] == 1627 # swap nibbles28 t_lo = tensor & 0x0F29 t_hi = tensor & 0xF030 t_swapped = (t_lo << 4) | (t_hi >> 4)31 tensor = t_swapped32 # transform aaaa...bbbb... to abababab...33 blk_a, blk_b = tensor.chunk(2, dim=-1)34 # get a_35 blk_a0 = (blk_a & 0xF0).view(-1, 1)36 blk_a1 = (blk_a << 4).view(-1, 1)37 blk_a = torch.stack((blk_a0, blk_a1), dim=2).view(tensor.shape)38 # get _b39 blk_b0 = (blk_b >> 4).view(-1, 1)40 blk_b1 = (blk_b & 0x0F).view(-1, 1)41 blk_b = torch.stack((blk_b0, blk_b1), dim=2).view(tensor.shape)42 # swap once more43 out = blk_a | blk_b44 out_h = out & 0xF045 out_l = out & 0x0F46 out = (out_h >> 4) | (out_l << 4)47 return out48 49 def repack_mxfp4(self, new_name: str, blocks: Tensor, scales: Tensor):50 assert blocks.dtype == torch.uint851 assert scales.dtype == torch.uint852 scales = scales.unsqueeze(-1)53 assert len(blocks.shape) == 454 assert len(scales.shape) == 455 blocks = self.transform_nibble_layout(blocks)56 new_data = torch.concat((scales, blocks), dim=-1)57 new_shape = [new_data.shape[0], new_data.shape[1], new_data.shape[2] * 32]58 logger.info(f"Repacked {new_name} with shape {new_shape} and quantization MXFP4")59 # flatten last dim60 new_data = new_data.view(new_data.shape[0], new_data.shape[1], new_data.shape[2] * new_data.shape[3])61 new_data = new_data.numpy()62 self.gguf_writer.add_tensor(new_name, new_data, raw_dtype=gguf.GGMLQuantizationType.MXFP4)63 64 def generate_extra_tensors(self) -> Iterable[tuple[str, Tensor]]:65 blocks0: Tensor = torch.zeros(1)66 blocks1: Tensor = torch.zeros(1)67 # we assume that tensors are loaded in the correct order68 for name, data_torch in self.get_tensors():69 if "mlp.experts.down_proj_blocks" in name:70 blocks0 = data_torch71 elif "mlp.experts.down_proj_scales" in name:72 new_name = self.map_tensor_name(name.replace("_scales", ".weight"))73 self.repack_mxfp4(new_name, blocks0, data_torch)74 elif "mlp.experts.gate_up_proj_blocks" in name:75 blocks0, blocks1 = data_torch[:, ::2, :, :], data_torch[:, 1::2, :, :]76 elif "mlp.experts.gate_up_proj_scales" in name:77 scales0, scales1 = data_torch[:, ::2, :], data_torch[:, 1::2, :]78 new_name_gate = self.map_tensor_name(name.replace("gate_up_proj_scales", "gate_proj.weight"))79 new_name_up = self.map_tensor_name(name.replace("gate_up_proj_scales", "up_proj.weight"))80 self.repack_mxfp4(new_name_gate, blocks0, scales0)81 self.repack_mxfp4(new_name_up, blocks1, scales1)82 return []83 84 @classmethod85 def filter_tensors(cls, item: tuple[str, Callable[[], Tensor]]) -> tuple[str, Callable[[], Tensor]] | None:86 name, gen = item87 88 if "sinks" in name:89 name += ".weight"90 91 return super().filter_tensors((name, gen))92 93 def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:94 # correct naming for down_proj95 if "down_proj" in name:96 if name.endswith("_bias"):97 name = name.replace("down_proj_bias", "down_proj.bias")98 elif "_blocks" not in name and "_scales" not in name:99 logger.warning(f"{name} is not in MXFP4, performance may be degraded")100 name = name.replace("down_proj", "down_proj.weight")101 data_torch = data_torch.transpose(-1, -2)102 else:103 # otherwise, it should already be repacked to ggml MXFP4 format104 return105 106 # split the gate_up into gate and up107 if "gate_up_proj" in name:108 if name.endswith("_bias"):109 name_up = name.replace("gate_up_proj_bias", "up_proj.bias")110 name_gate = name.replace("gate_up_proj_bias", "gate_proj.bias")111 gate_proj_bias, up_proj_bias = data_torch[..., ::2], data_torch[..., 1::2]112 yield from super().modify_tensors(gate_proj_bias, name_gate, bid)113 yield from super().modify_tensors(up_proj_bias, name_up, bid)114 elif "_blocks" not in name and "_scales" not in name:115 logger.warning(f"{name} is not in MXFP4, performance may be degraded")116 name_up = name.replace("gate_up_proj", "up_proj.weight")117 name_gate = name.replace("gate_up_proj", "gate_proj.weight")118 data_torch = data_torch.transpose(-1, -2)119 gate_proj_weight, up_proj_weight = data_torch[:, ::2, :], data_torch[:, 1::2, :]120 yield from super().modify_tensors(gate_proj_weight, name_gate, bid)121 yield from super().modify_tensors(up_proj_weight, name_up, bid)122 else:123 yield from super().modify_tensors(data_torch, name, bid)124 125 def set_vocab(self):126 self._set_vocab_gpt2()127 128 def set_gguf_parameters(self):129 super().set_gguf_parameters()130 self.gguf_writer.add_sliding_window(self.hparams["sliding_window"])131 self.gguf_writer.add_expert_feed_forward_length(self.hparams["intermediate_size"])132 