MCPcopy Create free account
hub / github.com/microsoft/BitNet / BitnetModel

Class BitnetModel

utils/convert-hf-to-gguf-bitnet.py:956–1081  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

954
955@Model.register("BitnetForCausalLM")
956class BitnetModel(Model):
957 model_arch = gguf.MODEL_ARCH.BITNET
958
959 def set_vocab(self):
960 self._set_vocab_sentencepiece()
961
962 def set_gguf_parameters(self):
963 super().set_gguf_parameters()
964
965 self.gguf_writer.add_vocab_size(self.hparams["vocab_size"])
966
967 self.gguf_writer.add_rope_scaling_type(gguf.RopeScalingType.LINEAR)
968 self.gguf_writer.add_rope_scaling_factor(1.0)
969
970 def weight_quant(self, weight):
971 dtype = weight.dtype
972 weight = weight.float()
973 s = 1 / weight.abs().mean().clamp(min=1e-5)
974 result = (weight * s).round().clamp(-1, 1) / s
975 return result.type(dtype)
976
977 def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:
978 # quant weight to i2 (in fp16)
979 if name.endswith(("q_proj.weight", "k_proj.weight", "v_proj.weight",
980 "down_proj.weight", "up_proj.weight", "gate_proj.weight",
981 "o_proj.weight")):
982 data_torch = self.weight_quant(data_torch)
983
984 return [(self.map_tensor_name(name), data_torch)]
985
986 def write_tensors(self):
987 max_name_len = max(len(s) for _, s in self.tensor_map.mapping.values()) + len(".weight,")
988
989 for name, data_torch in self.get_tensors():
990 # we don't need these
991 if name.endswith((".attention.masked_bias", ".attention.bias", ".rotary_emb.inv_freq")):
992 continue
993
994 old_dtype = data_torch.dtype
995
996 # convert any unsupported data types to float32
997 if data_torch.dtype not in (torch.float16, torch.float32):
998 data_torch = data_torch.to(torch.float32)
999
1000 # use the first number-like part of the tensor name as the block id
1001 bid = None
1002 for part in name.split("."):
1003 if part.isdecimal():
1004 bid = int(part)
1005 break
1006
1007 for new_name, data in ((n, d.squeeze().numpy()) for n, d in self.modify_tensors(data_torch, name, bid)):
1008 data: np.ndarray = data # type hint
1009 data_shape = data.shape
1010 n_dims = len(data.shape)
1011 data_dtype = data.dtype
1012 data_qtype: gguf.GGMLQuantizationType | None = None
1013

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected