File size: 6,263 Bytes
35cdf53 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 |
"""Model param loading."""
import bisect
import collections
from collections.abc import Iterator
import contextlib
import io
import os
import pathlib
import re
import struct
import sys
from typing import IO
import haiku as hk
import jax.numpy as jnp
import numpy as np
import zstandard
class RecordError(Exception):
"""Error reading a record."""
def encode_record(scope: str, name: str, arr: np.ndarray) -> bytes:
"""Encodes a single haiku param as bytes, preserving non-numpy dtypes."""
scope = scope.encode('utf-8')
name = name.encode('utf-8')
shape = arr.shape
dtype = str(arr.dtype).encode('utf-8')
arr = np.ascontiguousarray(arr)
if sys.byteorder == 'big':
arr = arr.byteswap()
arr_buffer = arr.tobytes('C')
header = struct.pack(
'<5i', len(scope), len(name), len(dtype), len(shape), len(arr_buffer)
)
return header + b''.join(
(scope, name, dtype, struct.pack(f'{len(shape)}i', *shape), arr_buffer)
)
def _read_record(stream: IO[bytes]) -> tuple[str, str, np.ndarray] | None:
"""Reads a record encoded by `_encode_record` from a byte stream."""
header_size = struct.calcsize('<5i')
header = stream.read(header_size)
if not header:
return None
if len(header) < header_size:
raise RecordError(f'Incomplete header: {len(header)=} < {header_size=}')
(scope_len, name_len, dtype_len, shape_len, arr_buffer_len) = struct.unpack(
'<5i', header
)
fmt = f'<{scope_len}s{name_len}s{dtype_len}s{shape_len}i'
payload_size = struct.calcsize(fmt) + arr_buffer_len
payload = stream.read(payload_size)
if len(payload) < payload_size:
raise RecordError(f'Incomplete payload: {len(payload)=} < {payload_size=}')
scope, name, dtype, *shape = struct.unpack_from(fmt, payload)
scope = scope.decode('utf-8')
name = name.decode('utf-8')
dtype = dtype.decode('utf-8')
arr = np.frombuffer(payload[-arr_buffer_len:], dtype=dtype)
arr = np.reshape(arr, shape)
if sys.byteorder == 'big':
arr = arr.byteswap()
return scope, name, arr
def read_records(stream: IO[bytes]) -> Iterator[tuple[str, str, np.ndarray]]:
"""Fully reads the contents of a byte stream."""
while record := _read_record(stream):
yield record
class _MultiFileIO(io.RawIOBase):
"""A file-like object that presents a concatenated view of multiple files."""
def __init__(self, files: list[pathlib.Path]):
self._files = files
self._stack = contextlib.ExitStack()
self._handles = [
self._stack.enter_context(file.open('rb')) for file in files
]
self._sizes = []
for handle in self._handles:
handle.seek(0, os.SEEK_END)
self._sizes.append(handle.tell())
self._length = sum(self._sizes)
self._offsets = [0]
for s in self._sizes[:-1]:
self._offsets.append(self._offsets[-1] + s)
self._abspos = 0
self._relpos = (0, 0)
def _abs_to_rel(self, pos: int) -> tuple[int, int]:
idx = bisect.bisect_right(self._offsets, pos) - 1
return idx, pos - self._offsets[idx]
def close(self):
self._stack.close()
def closed(self) -> bool:
return all(handle.closed for handle in self._handles)
def fileno(self) -> int:
return -1
def readable(self) -> bool:
return True
def tell(self) -> int:
return self._abspos
def seek(self, pos: int, whence: int = os.SEEK_SET, /):
match whence:
case os.SEEK_SET:
pass
case os.SEEK_CUR:
pos += self._abspos
case os.SEEK_END:
pos = self._length - pos
case _:
raise ValueError(f'Invalid whence: {whence}')
self._abspos = pos
self._relpos = self._abs_to_rel(pos)
def readinto(self, b: bytearray | memoryview) -> int:
result = 0
mem = memoryview(b)
while mem:
self._handles[self._relpos[0]].seek(self._relpos[1])
count = self._handles[self._relpos[0]].readinto(mem)
result += count
self._abspos += count
self._relpos = self._abs_to_rel(self._abspos)
mem = mem[count:]
if self._abspos == self._length:
break
return result
@contextlib.contextmanager
def open_for_reading(model_files: list[pathlib.Path], is_compressed: bool):
with contextlib.closing(_MultiFileIO(model_files)) as f:
if is_compressed:
yield zstandard.ZstdDecompressor().stream_reader(f)
else:
yield f
def _match_model(
paths: list[pathlib.Path], pattern: re.Pattern[str]
) -> dict[str, list[pathlib.Path]]:
"""Match files in a directory with a pattern, and group by model name."""
models = collections.defaultdict(list)
for path in paths:
match = pattern.fullmatch(path.name)
if match:
models[match.group('model_name')].append(path)
return {k: sorted(v) for k, v in models.items()}
def select_model_files(
model_dir: pathlib.Path, model_name: str | None = None
) -> tuple[list[pathlib.Path], bool]:
"""Select the model files from a model directory."""
files = [file for file in model_dir.iterdir() if file.is_file()]
for pattern, is_compressed in (
(r'(?P<model_name>.*)\.[0-9]+\.bin\.zst$', True),
(r'(?P<model_name>.*)\.bin\.zst\.[0-9]+$', True),
(r'(?P<model_name>.*)\.[0-9]+\.bin$', False),
(r'(?P<model_name>.*)\.bin]\.[0-9]+$', False),
(r'(?P<model_name>.*)\.bin\.zst$', True),
(r'(?P<model_name>.*)\.bin$', False),
):
models = _match_model(files, re.compile(pattern))
if model_name is not None:
if model_name in models:
return models[model_name], is_compressed
else:
if models:
if len(models) > 1:
raise RuntimeError(f'Multiple models matched in {model_dir}')
_, model_files = models.popitem()
return model_files, is_compressed
raise FileNotFoundError(f'No models matched in {model_dir}')
def get_model_haiku_params(model_dir: pathlib.Path) -> hk.Params:
"""Get the Haiku parameters from a model name."""
params: dict[str, dict[str, jnp.Array]] = {}
model_files, is_compressed = select_model_files(model_dir)
with open_for_reading(model_files, is_compressed) as stream:
for scope, name, arr in read_records(stream):
params.setdefault(scope, {})[name] = jnp.array(arr)
if not params:
raise FileNotFoundError(f'Model missing from "{model_dir}"')
return params
|