Compare commits

...

2 Commits

Author SHA1 Message Date
M. Yusuf Sarıgöz af1c9966c8 gguf : start write tensor info 2023-07-27 10:32:31 +03:00
M. Yusuf Sarıgöz 8332d26123 refactor: reduce code duplication and better API 2023-07-27 09:48:08 +03:00
2 changed files with 48 additions and 41 deletions
+1
View File
@@ -1,5 +1,6 @@
GGUF_MAGIC = 0x47475546
GGUF_VERSION = 1
GGUF_DEFAULT_ALIGNMENT = 32
# general
KEY_GENERAL_ARCHITECTURE = "general.architecture"
+47 -41
View File
@@ -7,7 +7,7 @@
import struct
from enum import IntEnum
from typing import List, Any
from typing import List, Any, Sequence
import constants
@@ -57,77 +57,71 @@ class GGUFValueType(IntEnum):
class GGUFWriter:
def __init__(self, buffered_writer):
self.buffered_writer = buffered_writer
def __init__(self, fout):
self.fout = fout
self.offset_tensor = 0
def write_header(self, tensor_count: int, metadata_kv_count: int):
self.buffered_writer.write(struct.pack("<I", constants.GGUF_MAGIC))
self.buffered_writer.write(struct.pack("<I", constants.GGUF_VERSION))
self.buffered_writer.write(struct.pack("<I", tensor_count))
self.buffered_writer.write(struct.pack("<I", metadata_kv_count))
self.fout.write(struct.pack("<I", constants.GGUF_MAGIC))
self.fout.write(struct.pack("<I", constants.GGUF_VERSION))
self.fout.write(struct.pack("<I", tensor_count))
self.fout.write(struct.pack("<I", metadata_kv_count))
@classmethod
def open(cls, path: str) -> "GGUFWriter":
f = open(path, "wb")
return cls(f)
def write_key(self, key: str, value_type: GGUFValueType):
encoded_key = key.encode("utf8")
self.buffered_writer.write(struct.pack("<I", len(encoded_key)))
self.buffered_writer.write(encoded_key)
self.buffered_writer.write(struct.pack("<I", value_type))
def write_key(self, key: str):
self.write_value(key, GGUFValueType.STRING)
def write_uint8(self, key: str, value: int):
self.write_key(key, GGUFValueType.UINT8)
self.buffered_writer.write(struct.pack("<B", value))
self.write_key(key)
self.write_value(value, GGUFValueType.UINT8)
def write_int8(self, key: str, value: int):
self.write_key(key, GGUFValueType.INT8)
self.buffered_writer.write(struct.pack("<b", value))
self.write_key(key)
self.write_value(value, GGUFValueType.INT8)
def write_uint16(self, key: str, value: int):
self.write_key(key, GGUFValueType.UINT16)
self.buffered_writer.write(struct.pack("<H", value))
self.write_key(key)
self.write_value(value, GGUFValueType.UINT16)
def write_int16(self, key: str, value: int):
self.write_key(key, GGUFValueType.INT16)
self.buffered_writer.write(struct.pack("<h", value))
self.write_key(key)
self.write_value(value, GGUFValueType.INT16)
def write_uint32(self, key: str, value: int):
self.write_key(key, GGUFValueType.UINT32)
self.buffered_writer.write(struct.pack("<I", value))
self.write_key(key)
self.write(value, GGUFValueType.UINT32)
def write_int32(self, key: str, value: int):
self.write_key(key, GGUFValueType.INT32)
self.buffered_writer.write(struct.pack("<i", value))
self.write_key(key)
self.write_value(value, GGUFValueType.INT32)
def write_float32(self, key: str, value: float):
self.write_key(key, GGUFValueType.FLOAT32)
self.buffered_writer.write(struct.pack("<f", value))
self.write_key(key)
self.write_value(value, GGUFValueType.FLOAT32)
def write_bool(self, key: str, value: bool):
self.write_key(key, GGUFValueType.BOOL)
self.buffered_writer.write(struct.pack("<?", value))
self.write_key(key)
self.write_value(value, GGUFValueType.BOOL)
def write_string(self, key: str, value: str):
self.write_key(key, GGUFValueType.STRING)
encoded_string = value.encode('utf-8')
self.buffered_writer.write(struct.pack("<I", len(encoded_string)))
self.buffered_writer.write(encoded_string)
self.write_key(key)
self.write_value(value, GGUFValueType.STRING)
def write_array(self, key: str, value: list):
if not isinstance(value, list):
raise ValueError("Value must be a list for array type")
self.write_key(key, GGUFValueType.ARRAY)
self.write_key(key)
self.write_value(value, GGUFValueType.ARRAY)
self.buffered_writer.write(struct.pack("<I", len(value)))
def write_value(self: str, value: Any, value_type: GGUFValueType = None):
if value_type is None:
value_type = GGUFValueType.get_type(value)
for item in value:
self.write_value(item)
def write_value(self: str, value: Any):
value_type = GGUFValueType.get_type(value)
self.buffered_writer.write(struct.pack("<I", value_type))
if value_type == GGUFValueType.UINT8:
@@ -157,11 +151,23 @@ class GGUFWriter:
else:
raise ValueError("Invalid GGUF metadata value type")
def write_tensor_info(self, name: str, shape: Sequence[int], dtype: GGMLQuantizationType):
self.write_value(name, GGUFValueType.STRING)
n_dims = len(shape)
self.write_value(n_dims, GGUFValueType.INT32)
for i in range(n_dims):
self.write_value(shape[n_dims - 1 - i], GGUFValueType.INT32)
self.fout.write(struct.pack("<Q", self.offset_tensor))
# TODO: update offset with alignment
# probably we need a dict as a class attribute to hold tensor data while writing
def flush(self):
self.buffered_writer.flush()
self.fout.flush()
def close(self):
self.buffered_writer.close()
self.fout.close()
def write_architecture(self, architecture: str):
self.write_string(constants.KEY_GENERAL_ARCHITECTURE,