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

Class Model

utils/generate-dummy-bitnet-model.py:120–494  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

118
119
120class Model(ABC):
121 _model_classes: dict[str, type[Model]] = {}
122
123 def __init__(self, dir_model: Path, ftype: int, fname_out: Path, is_big_endian: bool, use_temp_file: bool):
124 self.dir_model = dir_model
125 self.ftype = ftype
126 self.fname_out = fname_out
127 self.is_big_endian = is_big_endian
128 self.endianess = gguf.GGUFEndian.BIG if is_big_endian else gguf.GGUFEndian.LITTLE
129 self.use_temp_file = use_temp_file
130 self.is_safetensors = self._is_model_safetensors()
131 self.num_parts = Model.count_model_parts(self.dir_model, ".safetensors" if self.is_safetensors else ".bin")
132 self.part_names = self._get_part_names()
133 self.hparams = Model.load_hparams(self.dir_model)
134 self.gguf_writer = gguf.GGUFWriter(fname_out, gguf.MODEL_ARCH_NAMES[self.model_arch], endianess=self.endianess, use_temp_file=self.use_temp_file)
135 self.block_count = self.find_hparam(["n_layers", "num_hidden_layers", "n_layer"])
136 self.tensor_map = gguf.get_tensor_name_map(self.model_arch, self.block_count)
137
138 @property
139 @abstractmethod
140 def model_arch(self) -> gguf.MODEL_ARCH:
141 pass
142
143 def find_hparam(self, keys: Sequence[str], optional: bool = False) -> Any:
144 key = next((k for k in keys if k in self.hparams), None)
145 if key is not None:
146 return self.hparams[key]
147 if optional:
148 return None
149 raise KeyError(f"could not find any of: {keys}")
150
151 def set_vocab(self):
152 self._set_vocab_gpt2()
153
154 def get_tensors(self) -> Iterator[tuple[str, Tensor]]:
155 for part_name in self.part_names:
156 logger.info(f"gguf: loading model part '{part_name}'")
157 ctx: ContextManager[Any]
158 if self.is_safetensors:
159 from safetensors import safe_open
160 ctx = cast(ContextManager[Any], safe_open(self.dir_model / part_name, framework="pt", device="cpu"))
161 else:
162 ctx = contextlib.nullcontext(torch.load(str(self.dir_model / part_name), map_location="cpu", mmap=True, weights_only=True))
163
164 with ctx as model_part:
165 for name in model_part.keys():
166 data = model_part.get_tensor(name) if self.is_safetensors else model_part[name]
167 yield name, data
168
169 def match_model_tensor_name(self, name: str, key: gguf.MODEL_TENSOR, bid: int | None, suffix: str = ".weight") -> bool:
170 if key not in gguf.MODEL_TENSORS[self.model_arch]:
171 return False
172 key_name: str = gguf.TENSOR_NAMES[key]
173 if "{bid}" in key_name:
174 if bid is None:
175 return False
176 key_name = key_name.format(bid=bid)
177 else:

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected