mlx-dfloat 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- mlx_dfloat/__init__.py +25 -0
- mlx_dfloat/_memory_caps.py +74 -0
- mlx_dfloat/_metal_decode.py +420 -0
- mlx_dfloat/_safetensors.py +185 -0
- mlx_dfloat/_scrub.py +29 -0
- mlx_dfloat/_version.py +24 -0
- mlx_dfloat/_watchdog.py +253 -0
- mlx_dfloat/bench/__init__.py +4 -0
- mlx_dfloat/bench/capped.py +165 -0
- mlx_dfloat/bench/preflight.py +216 -0
- mlx_dfloat/bench/results.py +227 -0
- mlx_dfloat/bench/scenario.py +191 -0
- mlx_dfloat/bench/table.py +251 -0
- mlx_dfloat/cli.py +30 -0
- mlx_dfloat/decode.py +120 -0
- mlx_dfloat/errors.py +45 -0
- mlx_dfloat/format.py +462 -0
- mlx_dfloat/integrate/__init__.py +1 -0
- mlx_dfloat/integrate/coverage.py +122 -0
- mlx_dfloat/integrate/memory.py +55 -0
- mlx_dfloat/integrate/names.py +121 -0
- mlx_dfloat/integrate/placeholders.py +72 -0
- mlx_dfloat/integrate/providers.py +271 -0
- mlx_dfloat/integrate/seam.py +196 -0
- mlx_dfloat/mflux/__init__.py +31 -0
- mlx_dfloat/mflux/flux1/__init__.py +1 -0
- mlx_dfloat/mflux/flux1/cli.py +419 -0
- mlx_dfloat/mflux/flux1/init.py +245 -0
- mlx_dfloat/mflux/flux1/lifecycle.py +159 -0
- mlx_dfloat/mflux/flux1/memory.py +201 -0
- mlx_dfloat/mflux/flux1/model.py +553 -0
- mlx_dfloat/mflux/flux1/names.py +79 -0
- mlx_dfloat/mflux/flux1/transformer.py +240 -0
- mlx_dfloat/py.typed +0 -0
- mlx_dfloat/reference.py +241 -0
- mlx_dfloat-0.1.0.dist-info/METADATA +325 -0
- mlx_dfloat-0.1.0.dist-info/RECORD +41 -0
- mlx_dfloat-0.1.0.dist-info/WHEEL +4 -0
- mlx_dfloat-0.1.0.dist-info/entry_points.txt +2 -0
- mlx_dfloat-0.1.0.dist-info/licenses/LICENSE +202 -0
- mlx_dfloat-0.1.0.dist-info/licenses/NOTICE +40 -0
|
@@ -0,0 +1,185 @@
|
|
|
1
|
+
"""Minimal safetensors reader: header parsing plus lazy, read-only memmap access.
|
|
2
|
+
|
|
3
|
+
Written against the safetensors format (8-byte little-endian header length, JSON header, raw
|
|
4
|
+
little-endian data). Files come from the network, so every field is validated and the header size
|
|
5
|
+
is capped. BF16 tensors are returned as their uint16 bit patterns.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
import json
|
|
9
|
+
import math
|
|
10
|
+
import struct
|
|
11
|
+
from collections.abc import Mapping
|
|
12
|
+
from dataclasses import dataclass
|
|
13
|
+
from pathlib import Path
|
|
14
|
+
from typing import Any
|
|
15
|
+
|
|
16
|
+
import numpy as np
|
|
17
|
+
import numpy.typing as npt
|
|
18
|
+
|
|
19
|
+
from mlx_dfloat.errors import DFloatFormatError
|
|
20
|
+
|
|
21
|
+
MAX_HEADER_BYTES = 100_000_000
|
|
22
|
+
_MAX_REPR_CHARS = 80
|
|
23
|
+
|
|
24
|
+
DTYPES: Mapping[str, np.dtype[Any]] = {
|
|
25
|
+
"BOOL": np.dtype(np.bool_),
|
|
26
|
+
"U8": np.dtype("u1"),
|
|
27
|
+
"I8": np.dtype("i1"),
|
|
28
|
+
"U16": np.dtype("<u2"),
|
|
29
|
+
"I16": np.dtype("<i2"),
|
|
30
|
+
"F16": np.dtype("<f2"),
|
|
31
|
+
"BF16": np.dtype("<u2"),
|
|
32
|
+
"U32": np.dtype("<u4"),
|
|
33
|
+
"I32": np.dtype("<i4"),
|
|
34
|
+
"F32": np.dtype("<f4"),
|
|
35
|
+
"U64": np.dtype("<u8"),
|
|
36
|
+
"I64": np.dtype("<i8"),
|
|
37
|
+
"F64": np.dtype("<f8"),
|
|
38
|
+
}
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
@dataclass(frozen=True, slots=True, kw_only=True)
|
|
42
|
+
class TensorInfo:
|
|
43
|
+
"""Location and type of one tensor inside a safetensors file."""
|
|
44
|
+
|
|
45
|
+
name: str
|
|
46
|
+
dtype: str
|
|
47
|
+
shape: tuple[int, ...]
|
|
48
|
+
offset: int
|
|
49
|
+
nbytes: int
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def short_repr(value: object) -> str:
|
|
53
|
+
"""``repr(value)`` capped at 80 characters, for echoing untrusted header/config values."""
|
|
54
|
+
text = repr(value)
|
|
55
|
+
return text if len(text) <= _MAX_REPR_CHARS else text[: _MAX_REPR_CHARS - 3] + "..."
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
def _no_duplicates(pairs: list[tuple[str, Any]]) -> dict[str, Any]:
|
|
59
|
+
keys = [k for k, _ in pairs]
|
|
60
|
+
if len(keys) != len(set(keys)):
|
|
61
|
+
raise ValueError("duplicate key in safetensors header")
|
|
62
|
+
return dict(pairs)
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
def parse_header(
|
|
66
|
+
raw: bytes, *, data_start: int, file_size: int, source: str
|
|
67
|
+
) -> dict[str, TensorInfo]:
|
|
68
|
+
"""Parse and validate a safetensors JSON header.
|
|
69
|
+
|
|
70
|
+
Args:
|
|
71
|
+
raw: The JSON header bytes.
|
|
72
|
+
data_start: Absolute file offset where tensor data begins (8 + header length).
|
|
73
|
+
file_size: Total file size in bytes, used to detect truncated downloads.
|
|
74
|
+
source: File name used in error messages.
|
|
75
|
+
|
|
76
|
+
Returns:
|
|
77
|
+
Tensor name to TensorInfo, without the ``__metadata__`` entry.
|
|
78
|
+
|
|
79
|
+
Raises:
|
|
80
|
+
DFloatFormatError: The header is not valid JSON, has duplicate keys, or an entry is
|
|
81
|
+
malformed or out of bounds.
|
|
82
|
+
"""
|
|
83
|
+
try:
|
|
84
|
+
header = json.loads(raw, object_pairs_hook=_no_duplicates)
|
|
85
|
+
except ValueError as exc: # JSONDecodeError, UnicodeDecodeError, duplicate keys
|
|
86
|
+
message = "duplicate key" if "duplicate" in str(exc) else "not valid JSON"
|
|
87
|
+
raise DFloatFormatError(f"{source}: safetensors header is {message}") from exc
|
|
88
|
+
except RecursionError as exc:
|
|
89
|
+
raise DFloatFormatError(f"{source}: safetensors header is nested too deeply") from exc
|
|
90
|
+
if not isinstance(header, dict):
|
|
91
|
+
raise DFloatFormatError(f"{source}: safetensors header is not a JSON object")
|
|
92
|
+
infos: dict[str, TensorInfo] = {}
|
|
93
|
+
for name, meta in header.items():
|
|
94
|
+
if name == "__metadata__":
|
|
95
|
+
continue
|
|
96
|
+
infos[name] = _parse_entry(
|
|
97
|
+
name, meta, data_start=data_start, file_size=file_size, source=source
|
|
98
|
+
)
|
|
99
|
+
return infos
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
def _parse_entry(
|
|
103
|
+
name: str, meta: object, *, data_start: int, file_size: int, source: str
|
|
104
|
+
) -> TensorInfo:
|
|
105
|
+
label = short_repr(name)
|
|
106
|
+
if not isinstance(meta, dict) or {"dtype", "shape", "data_offsets"} - meta.keys():
|
|
107
|
+
raise DFloatFormatError(f"{source}: tensor entry {label} is malformed")
|
|
108
|
+
dtype = meta["dtype"]
|
|
109
|
+
if not isinstance(dtype, str) or dtype not in DTYPES:
|
|
110
|
+
raise DFloatFormatError(
|
|
111
|
+
f"{source}: tensor {label} has unsupported dtype {short_repr(dtype)}"
|
|
112
|
+
)
|
|
113
|
+
shape = meta["shape"]
|
|
114
|
+
# No real dimension exceeds the file size; the bound also keeps a zero-element shape such as
|
|
115
|
+
# [0, 2**63] (product 0, so the byte-size check passes) from reaching np.empty.
|
|
116
|
+
if not isinstance(shape, list) or not all(
|
|
117
|
+
type(d) is int and 0 <= d <= file_size for d in shape
|
|
118
|
+
):
|
|
119
|
+
raise DFloatFormatError(
|
|
120
|
+
f"{source}: tensor {label} has an invalid shape {short_repr(shape)}"
|
|
121
|
+
)
|
|
122
|
+
offsets = meta["data_offsets"]
|
|
123
|
+
if (
|
|
124
|
+
not isinstance(offsets, list)
|
|
125
|
+
or len(offsets) != 2
|
|
126
|
+
or not all(type(o) is int for o in offsets)
|
|
127
|
+
or not 0 <= offsets[0] <= offsets[1]
|
|
128
|
+
):
|
|
129
|
+
raise DFloatFormatError(
|
|
130
|
+
f"{source}: tensor {label} has invalid data offsets {short_repr(offsets)}"
|
|
131
|
+
)
|
|
132
|
+
nbytes = offsets[1] - offsets[0]
|
|
133
|
+
if nbytes != math.prod(shape) * DTYPES[dtype].itemsize:
|
|
134
|
+
raise DFloatFormatError(
|
|
135
|
+
f"{source}: tensor {label} byte size does not match its shape and dtype"
|
|
136
|
+
)
|
|
137
|
+
if data_start + offsets[1] > file_size:
|
|
138
|
+
raise DFloatFormatError(
|
|
139
|
+
f"{source}: tensor {label} extends past end of file ({file_size} bytes); partial download?"
|
|
140
|
+
)
|
|
141
|
+
return TensorInfo(
|
|
142
|
+
name=name, dtype=dtype, shape=tuple(shape), offset=data_start + offsets[0], nbytes=nbytes
|
|
143
|
+
)
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
def read_header(path: Path) -> dict[str, TensorInfo]:
|
|
147
|
+
"""Read and validate the header of a safetensors file without reading tensor data.
|
|
148
|
+
|
|
149
|
+
Raises:
|
|
150
|
+
DFloatFormatError: The file cannot be read, or its header is invalid or oversized.
|
|
151
|
+
"""
|
|
152
|
+
try:
|
|
153
|
+
file_size = path.stat().st_size
|
|
154
|
+
with path.open("rb") as handle:
|
|
155
|
+
prefix = handle.read(8)
|
|
156
|
+
if len(prefix) < 8:
|
|
157
|
+
raise DFloatFormatError(f"{path.name}: too short for a safetensors header")
|
|
158
|
+
(length,) = struct.unpack("<Q", prefix)
|
|
159
|
+
if length == 0 or length > MAX_HEADER_BYTES or 8 + length > file_size:
|
|
160
|
+
raise DFloatFormatError(
|
|
161
|
+
f"{path.name}: safetensors header length {length} is invalid"
|
|
162
|
+
)
|
|
163
|
+
raw = handle.read(length)
|
|
164
|
+
except OSError as exc:
|
|
165
|
+
raise DFloatFormatError(f"{path}: cannot read file ({exc.strerror or exc})") from exc
|
|
166
|
+
return parse_header(raw, data_start=8 + length, file_size=file_size, source=path.name)
|
|
167
|
+
|
|
168
|
+
|
|
169
|
+
def read_array(path: Path, info: TensorInfo) -> npt.NDArray[Any]:
|
|
170
|
+
"""Return a read-only, lazily paged view of one tensor (BF16 as uint16 bits)."""
|
|
171
|
+
dtype = DTYPES[info.dtype]
|
|
172
|
+
if info.nbytes == 0:
|
|
173
|
+
return np.empty(info.shape, dtype=dtype)
|
|
174
|
+
return np.memmap(path, dtype=dtype, mode="r", offset=info.offset, shape=info.shape)
|
|
175
|
+
|
|
176
|
+
|
|
177
|
+
__all__ = [
|
|
178
|
+
"DTYPES",
|
|
179
|
+
"MAX_HEADER_BYTES",
|
|
180
|
+
"TensorInfo",
|
|
181
|
+
"parse_header",
|
|
182
|
+
"read_array",
|
|
183
|
+
"read_header",
|
|
184
|
+
"short_repr",
|
|
185
|
+
]
|
mlx_dfloat/_scrub.py
ADDED
|
@@ -0,0 +1,29 @@
|
|
|
1
|
+
"""Home-directory scrubbing for JSON the tools write (result files are committed to a public repo)."""
|
|
2
|
+
|
|
3
|
+
import re
|
|
4
|
+
from collections.abc import Mapping
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
from typing import Any
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def scrub_home(value: Any, home: str = str(Path.home())) -> Any:
|
|
10
|
+
"""``value`` with every occurrence of ``home`` as a whole path component written as ``~``.
|
|
11
|
+
|
|
12
|
+
Walks dicts, lists and tuples (tuples come back as lists, as JSON writes them); a sibling
|
|
13
|
+
directory that merely shares the prefix (``/home/ab2`` for ``/home/ab``) is left alone.
|
|
14
|
+
"""
|
|
15
|
+
home = home.rstrip("/")
|
|
16
|
+
if not home:
|
|
17
|
+
return value
|
|
18
|
+
pattern = re.compile(re.escape(home) + r"(?=/|\s|$)")
|
|
19
|
+
return _scrub(value, pattern)
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def _scrub(value: Any, pattern: re.Pattern[str]) -> Any:
|
|
23
|
+
if isinstance(value, str):
|
|
24
|
+
return pattern.sub("~", value)
|
|
25
|
+
if isinstance(value, Mapping):
|
|
26
|
+
return {k: _scrub(v, pattern) for k, v in value.items()}
|
|
27
|
+
if isinstance(value, list | tuple):
|
|
28
|
+
return [_scrub(v, pattern) for v in value]
|
|
29
|
+
return value
|
mlx_dfloat/_version.py
ADDED
|
@@ -0,0 +1,24 @@
|
|
|
1
|
+
# file generated by vcs-versioning
|
|
2
|
+
# don't change, don't track in version control
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
__all__ = [
|
|
6
|
+
"__version__",
|
|
7
|
+
"__version_tuple__",
|
|
8
|
+
"version",
|
|
9
|
+
"version_tuple",
|
|
10
|
+
"__commit_id__",
|
|
11
|
+
"commit_id",
|
|
12
|
+
]
|
|
13
|
+
|
|
14
|
+
version: str
|
|
15
|
+
__version__: str
|
|
16
|
+
__version_tuple__: tuple[int | str, ...]
|
|
17
|
+
version_tuple: tuple[int | str, ...]
|
|
18
|
+
commit_id: str | None
|
|
19
|
+
__commit_id__: str | None
|
|
20
|
+
|
|
21
|
+
__version__ = version = '0.1.0'
|
|
22
|
+
__version_tuple__ = version_tuple = (0, 1, 0)
|
|
23
|
+
|
|
24
|
+
__commit_id__ = commit_id = None
|
mlx_dfloat/_watchdog.py
ADDED
|
@@ -0,0 +1,253 @@
|
|
|
1
|
+
"""Process-level watchdog for heavy scripts: a memory ceiling plus a wall-clock backstop.
|
|
2
|
+
|
|
3
|
+
The ceiling is checked against the larger of the process's OS-accounted memory footprint
|
|
4
|
+
(``phys_footprint``, what macOS memory pressure sees) and MLX active + cache memory, never a sum
|
|
5
|
+
of RSS and MLX: a buffer loaded with ``mx.load`` appears in both RSS and MLX active memory
|
|
6
|
+
(verified: a 1 GiB load moves both by about 1 GiB), so a sum double-counts it and would
|
|
7
|
+
false-abort at half the real ceiling, while the maximum still catches an overrun held in MLX's
|
|
8
|
+
cache pool. The abort artifact names the counter that tripped (``verdict_counter``) and records
|
|
9
|
+
the peaks of the footprint, of MLX active + cache, and of the watched maximum. ``rss``,
|
|
10
|
+
``mlx_active``, and ``mlx_cache`` ride along in every sample and abort artifact as diagnostics. A
|
|
11
|
+
sampling failure (psutil, MLX, the footprint read, or the artifact write itself) still aborts the
|
|
12
|
+
process instead of leaving the job running unwatched; its artifact names no counter
|
|
13
|
+
(``verdict_counter: "none"``, ``verdict_memory: null``). A caller may pass a ``context`` mapping
|
|
14
|
+
(``generate`` passes its model, size, seed and steps); it is written verbatim under ``"context"``
|
|
15
|
+
in the abort artifact, so the artifact names the run it stopped. Without one there is no
|
|
16
|
+
``"context"`` key. ``peak_mlx`` is None until a sample has read MLX's counters.
|
|
17
|
+
"""
|
|
18
|
+
|
|
19
|
+
import ctypes
|
|
20
|
+
import json
|
|
21
|
+
import os
|
|
22
|
+
import sys
|
|
23
|
+
import threading
|
|
24
|
+
import time
|
|
25
|
+
from collections.abc import Mapping
|
|
26
|
+
from pathlib import Path
|
|
27
|
+
from typing import Any
|
|
28
|
+
|
|
29
|
+
import mlx.core as mx
|
|
30
|
+
|
|
31
|
+
from mlx_dfloat.errors import DFloatDependencyError
|
|
32
|
+
|
|
33
|
+
EXIT_MEMORY = 70
|
|
34
|
+
EXIT_WALL = 71
|
|
35
|
+
|
|
36
|
+
# The one process exit the watchdog makes; tests patch this alias, never os._exit process-wide.
|
|
37
|
+
_exit = os._exit
|
|
38
|
+
|
|
39
|
+
_RUSAGE_INFO_V2 = 2
|
|
40
|
+
# Loaded once at import time, not per 0.05 s sample.
|
|
41
|
+
_LIBPROC = ctypes.CDLL("/usr/lib/libproc.dylib") if sys.platform == "darwin" else None
|
|
42
|
+
if _LIBPROC is not None:
|
|
43
|
+
_LIBPROC.proc_pid_rusage.argtypes = [ctypes.c_int, ctypes.c_int, ctypes.c_void_p]
|
|
44
|
+
_LIBPROC.proc_pid_rusage.restype = ctypes.c_int
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
class _RusageInfoV2(ctypes.Structure):
|
|
48
|
+
"""Mirrors macOS's ``rusage_info_v2`` (``proc_info.h``); only ``ri_phys_footprint`` is read."""
|
|
49
|
+
|
|
50
|
+
_fields_ = [("ri_uuid", ctypes.c_uint8 * 16)] + [
|
|
51
|
+
(name, ctypes.c_uint64)
|
|
52
|
+
for name in (
|
|
53
|
+
"ri_user_time",
|
|
54
|
+
"ri_system_time",
|
|
55
|
+
"ri_pkg_idle_wkups",
|
|
56
|
+
"ri_interrupt_wkups",
|
|
57
|
+
"ri_pageins",
|
|
58
|
+
"ri_wired_size",
|
|
59
|
+
"ri_resident_size",
|
|
60
|
+
"ri_phys_footprint",
|
|
61
|
+
"ri_proc_start_abstime",
|
|
62
|
+
"ri_proc_exit_abstime",
|
|
63
|
+
"ri_child_user_time",
|
|
64
|
+
"ri_child_system_time",
|
|
65
|
+
"ri_child_pkg_idle_wkups",
|
|
66
|
+
"ri_child_interrupt_wkups",
|
|
67
|
+
"ri_child_pageins",
|
|
68
|
+
"ri_child_elapsed_abstime",
|
|
69
|
+
"ri_diskio_bytesread",
|
|
70
|
+
"ri_diskio_byteswritten",
|
|
71
|
+
)
|
|
72
|
+
]
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def _psutil() -> Any:
|
|
76
|
+
"""The ``psutil`` module, or a package-rooted error naming it (it ships in the mflux extra)."""
|
|
77
|
+
try:
|
|
78
|
+
import psutil
|
|
79
|
+
except ImportError as exc:
|
|
80
|
+
raise DFloatDependencyError(
|
|
81
|
+
"the watchdog needs psutil: install mlx-dfloat[mflux] or `pip install psutil`"
|
|
82
|
+
) from exc
|
|
83
|
+
return psutil
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def phys_footprint() -> int:
|
|
87
|
+
"""The OS-accounted footprint of this process (what macOS memory pressure sees); RSS on other platforms."""
|
|
88
|
+
if _LIBPROC is None:
|
|
89
|
+
return int(_psutil().Process().memory_info().rss)
|
|
90
|
+
info = _RusageInfoV2()
|
|
91
|
+
if _LIBPROC.proc_pid_rusage(os.getpid(), _RUSAGE_INFO_V2, ctypes.byref(info)) != 0:
|
|
92
|
+
raise OSError("proc_pid_rusage failed")
|
|
93
|
+
return int(info.ri_phys_footprint)
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
def watched_memory(*, footprint: int, mlx_active: int, mlx_cache: int) -> tuple[int, str]:
|
|
97
|
+
"""The number the ceiling is enforced against and which counter produced it.
|
|
98
|
+
|
|
99
|
+
The number is max(OS footprint, MLX active + cache). A maximum cannot double count the way a
|
|
100
|
+
sum of RSS and MLX active did (an ``mx.load``ed array lands in both), and it catches an
|
|
101
|
+
overrun that sits in MLX's cache pool, which the footprint alone reports late. A tie names
|
|
102
|
+
the footprint.
|
|
103
|
+
"""
|
|
104
|
+
mlx = mlx_active + mlx_cache
|
|
105
|
+
return (footprint, "footprint") if footprint >= mlx else (mlx, "mlx")
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def verdict(*, memory: int, ceiling: int, elapsed: float, budget: float) -> str | None:
|
|
109
|
+
"""Decide whether to abort: "memory", "wall", or None. ``memory`` is ``watched_memory``'s number."""
|
|
110
|
+
if memory > ceiling:
|
|
111
|
+
return "memory"
|
|
112
|
+
if elapsed > budget:
|
|
113
|
+
return "wall"
|
|
114
|
+
return None
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
def default_ceiling() -> int:
|
|
118
|
+
"""Physical memory minus a 4 GiB reserve for the OS and other processes.
|
|
119
|
+
|
|
120
|
+
Raises:
|
|
121
|
+
DFloatDependencyError: ``psutil`` is not installed.
|
|
122
|
+
"""
|
|
123
|
+
return int(_psutil().virtual_memory().total) - 4 * 1024**3
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
class Watchdog:
|
|
127
|
+
"""Background sampler that aborts the process with an honest artifact."""
|
|
128
|
+
|
|
129
|
+
def __init__(
|
|
130
|
+
self,
|
|
131
|
+
out_dir: Path,
|
|
132
|
+
*,
|
|
133
|
+
ceiling: int,
|
|
134
|
+
budget: float,
|
|
135
|
+
interval: float = 0.05,
|
|
136
|
+
context: Mapping[str, Any] | None = None,
|
|
137
|
+
) -> None:
|
|
138
|
+
"""Configure the ceiling (bytes), wall budget (seconds), poll interval and run context.
|
|
139
|
+
|
|
140
|
+
``context`` (JSON-serialisable) is written verbatim under ``"context"`` in the abort
|
|
141
|
+
artifact; None writes no ``"context"`` key.
|
|
142
|
+
|
|
143
|
+
Raises:
|
|
144
|
+
DFloatDependencyError: ``psutil`` (the RSS diagnostic) is not installed; refused here,
|
|
145
|
+
before any sampling, so a missing module never reads as a sampling failure.
|
|
146
|
+
"""
|
|
147
|
+
self._psutil = _psutil()
|
|
148
|
+
self.out_dir, self.ceiling, self.budget, self.interval = out_dir, ceiling, budget, interval
|
|
149
|
+
self.context = None if context is None else dict(context)
|
|
150
|
+
# Peaks seen so far: the OS footprint, MLX active + cache (None until a sample read MLX's
|
|
151
|
+
# counters), and the watched maximum the ceiling is enforced on. Read and written under
|
|
152
|
+
# ``_lock``, so a reset cannot interleave with a sample's update.
|
|
153
|
+
self.peak_footprint = 0
|
|
154
|
+
self.peak_mlx: int | None = None
|
|
155
|
+
self.peak_watched = 0
|
|
156
|
+
self._stop = threading.Event()
|
|
157
|
+
# Held while an abort is written: stop() waits for it, so no abort can land after the
|
|
158
|
+
# caller has stopped the watchdog and written its own verdict.
|
|
159
|
+
self._lock = threading.Lock()
|
|
160
|
+
self._thread = threading.Thread(target=self._run, daemon=True)
|
|
161
|
+
self._start = time.monotonic()
|
|
162
|
+
|
|
163
|
+
def start(self) -> "Watchdog":
|
|
164
|
+
"""Start sampling."""
|
|
165
|
+
self._start = time.monotonic()
|
|
166
|
+
self._thread.start()
|
|
167
|
+
return self
|
|
168
|
+
|
|
169
|
+
def reset_peak(self) -> None:
|
|
170
|
+
"""Start all three peaks over, so a later window's peak is not hidden by an earlier spike."""
|
|
171
|
+
with self._lock:
|
|
172
|
+
self.peak_footprint = 0
|
|
173
|
+
self.peak_mlx = None
|
|
174
|
+
self.peak_watched = 0
|
|
175
|
+
|
|
176
|
+
def stop(self) -> None:
|
|
177
|
+
"""Stop sampling. After this returns the watchdog never writes an abort or exits."""
|
|
178
|
+
with self._lock:
|
|
179
|
+
self._stop.set()
|
|
180
|
+
if self._thread.is_alive():
|
|
181
|
+
self._thread.join(timeout=1)
|
|
182
|
+
|
|
183
|
+
def _sample(self) -> tuple[str | None, dict[str, float | str | None]]:
|
|
184
|
+
elapsed = time.monotonic() - self._start
|
|
185
|
+
# The verdict fields stay "none" / None until a number was compared with the ceiling, so a
|
|
186
|
+
# sample error's artifact does not read as a footprint verdict.
|
|
187
|
+
sample: dict[str, float | str | None] = {
|
|
188
|
+
"footprint": 0,
|
|
189
|
+
"rss": 0,
|
|
190
|
+
"mlx_active": 0,
|
|
191
|
+
"mlx_cache": 0,
|
|
192
|
+
"elapsed": elapsed,
|
|
193
|
+
"verdict_memory": None,
|
|
194
|
+
"verdict_counter": "none",
|
|
195
|
+
}
|
|
196
|
+
try:
|
|
197
|
+
footprint = int(phys_footprint())
|
|
198
|
+
sample["footprint"] = footprint
|
|
199
|
+
sample["rss"] = int(self._psutil.Process().memory_info().rss)
|
|
200
|
+
active = int(mx.get_active_memory())
|
|
201
|
+
cache = int(mx.get_cache_memory())
|
|
202
|
+
sample["mlx_active"] = active
|
|
203
|
+
sample["mlx_cache"] = cache
|
|
204
|
+
memory, counter = watched_memory(
|
|
205
|
+
footprint=footprint, mlx_active=active, mlx_cache=cache
|
|
206
|
+
)
|
|
207
|
+
sample["verdict_memory"] = memory
|
|
208
|
+
sample["verdict_counter"] = counter
|
|
209
|
+
with self._lock:
|
|
210
|
+
self.peak_footprint = max(self.peak_footprint, footprint)
|
|
211
|
+
self.peak_mlx = max(self.peak_mlx or 0, active + cache)
|
|
212
|
+
self.peak_watched = max(self.peak_watched, memory)
|
|
213
|
+
reason = verdict(
|
|
214
|
+
memory=memory, ceiling=self.ceiling, elapsed=elapsed, budget=self.budget
|
|
215
|
+
)
|
|
216
|
+
except Exception: # a dead sampler must still abort, not run the job unwatched
|
|
217
|
+
reason = "sample_error"
|
|
218
|
+
return reason, sample
|
|
219
|
+
|
|
220
|
+
def _run(self) -> None:
|
|
221
|
+
while not self._stop.wait(self.interval):
|
|
222
|
+
reason, sample = self._sample()
|
|
223
|
+
if reason is not None:
|
|
224
|
+
self._fire(reason, sample)
|
|
225
|
+
return
|
|
226
|
+
|
|
227
|
+
def _fire(self, reason: str, sample: dict[str, float | str | None]) -> None:
|
|
228
|
+
with self._lock:
|
|
229
|
+
if self._stop.is_set():
|
|
230
|
+
return # the caller already stopped us and recorded its own result
|
|
231
|
+
code = EXIT_WALL if reason == "wall" else EXIT_MEMORY
|
|
232
|
+
try:
|
|
233
|
+
self.out_dir.mkdir(parents=True, exist_ok=True)
|
|
234
|
+
artifact: dict[str, Any] = {
|
|
235
|
+
"reason": reason,
|
|
236
|
+
"footprint": sample["footprint"],
|
|
237
|
+
"peak_footprint": self.peak_footprint,
|
|
238
|
+
"ceiling": self.ceiling,
|
|
239
|
+
"elapsed": sample["elapsed"],
|
|
240
|
+
"budget": self.budget,
|
|
241
|
+
"rss": sample["rss"],
|
|
242
|
+
"mlx_active": sample["mlx_active"],
|
|
243
|
+
"mlx_cache": sample["mlx_cache"],
|
|
244
|
+
"verdict_memory": sample.get("verdict_memory"),
|
|
245
|
+
"verdict_counter": sample.get("verdict_counter", "none"),
|
|
246
|
+
"peak_watched": self.peak_watched,
|
|
247
|
+
"peak_mlx": self.peak_mlx,
|
|
248
|
+
}
|
|
249
|
+
if self.context is not None:
|
|
250
|
+
artifact["context"] = self.context
|
|
251
|
+
(self.out_dir / "abort.json").write_text(json.dumps(artifact, indent=1))
|
|
252
|
+
finally:
|
|
253
|
+
_exit(code) # always exits, even if the artifact write above raised
|
|
@@ -0,0 +1,165 @@
|
|
|
1
|
+
"""Capped mode: emulate a smaller Mac's MLX limits on this one.
|
|
2
|
+
|
|
3
|
+
A real Mac's MLX defaults are ``memory_limit = cache_limit = min(1.5 x recommended working set,
|
|
4
|
+
0.95 x RAM)``; capped mode installs that memory limit for the tier, so MLX throttles no earlier than
|
|
5
|
+
it would on the real machine. MLX starts reclaiming its buffer cache at ``min(memory_limit, 0.95 x
|
|
6
|
+
recommended)`` (verified 2026-09-29 on mlx 0.32.2), so that is the cache limit installed as the
|
|
7
|
+
stand-in for the real reclaim point; the worker then installs its own, smaller cache limit and
|
|
8
|
+
records the effective value. The watchdog ceiling is the tier's budget (its recommended working
|
|
9
|
+
set) minus a reserve; crossing it aborts the run, because ``set_memory_limit`` only throttles and
|
|
10
|
+
never fails an allocation. The wired limit is 0 (not part of the emulation). The recommended
|
|
11
|
+
working set comes from device data when a datapoint exists, else from the 2/3 (16 and 24 GB) and
|
|
12
|
+
3/4 (above 24 GB) ratios, and for the host tier from the host's own value; the host tier keeps the
|
|
13
|
+
host caps (``is_host``), so MEASURED rows run under the limits a plain ``generate`` uses.
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
import dataclasses
|
|
17
|
+
|
|
18
|
+
import mlx.core as mx
|
|
19
|
+
|
|
20
|
+
from mlx_dfloat.errors import DFloatUnsupportedError
|
|
21
|
+
|
|
22
|
+
GIB = 1024**3
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
@dataclasses.dataclass(frozen=True, slots=True, kw_only=True)
|
|
26
|
+
class TierLimits:
|
|
27
|
+
"""The limits capped mode installs for one tier, and the label its rows get."""
|
|
28
|
+
|
|
29
|
+
tier_gb: int
|
|
30
|
+
ram_bytes: int
|
|
31
|
+
recommended_bytes: int
|
|
32
|
+
memory_limit_bytes: int
|
|
33
|
+
cache_limit_bytes: int
|
|
34
|
+
wired_limit_bytes: int
|
|
35
|
+
reserve_bytes: int
|
|
36
|
+
budget_bytes: int
|
|
37
|
+
ceiling_bytes: int
|
|
38
|
+
budget_source: str
|
|
39
|
+
is_host: bool
|
|
40
|
+
label: str
|
|
41
|
+
|
|
42
|
+
def as_dict(self) -> dict[str, int | str | bool]:
|
|
43
|
+
"""Every field, for a result file."""
|
|
44
|
+
return dataclasses.asdict(self)
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def reserve_for(tier_gb: int) -> int:
|
|
48
|
+
"""1.5 GiB up to 24 GB, 2 GiB above (initial values)."""
|
|
49
|
+
return int(1.5 * GIB) if tier_gb <= 24 else 2 * GIB
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def recommended_ratio(tier_gb: int) -> float:
|
|
53
|
+
"""The share of RAM assumed as the recommended working set until device data exists: 2/3 up to 24 GB, 0.75 above 24."""
|
|
54
|
+
return 2 / 3 if tier_gb <= 24 else 0.75
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def host_tier_gb(host_ram_bytes: int) -> int:
|
|
58
|
+
"""The host's own tier: its RAM rounded to whole GB."""
|
|
59
|
+
return round(host_ram_bytes / GIB)
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def tier_limits(
|
|
63
|
+
tier_gb: int,
|
|
64
|
+
*,
|
|
65
|
+
host_ram_bytes: int,
|
|
66
|
+
host_recommended_bytes: int,
|
|
67
|
+
device_recommended_bytes: int | None = None,
|
|
68
|
+
) -> TierLimits:
|
|
69
|
+
"""The limits for ``tier_gb`` on a host with ``host_ram_bytes`` and ``host_recommended_bytes``.
|
|
70
|
+
|
|
71
|
+
Raises:
|
|
72
|
+
DFloatUnsupportedError: The tier is larger than the host (it cannot be emulated here).
|
|
73
|
+
"""
|
|
74
|
+
host_tier = host_tier_gb(host_ram_bytes)
|
|
75
|
+
if tier_gb > host_tier:
|
|
76
|
+
raise DFloatUnsupportedError(
|
|
77
|
+
f"a {tier_gb} GB tier cannot be emulated on a {host_tier} GB host"
|
|
78
|
+
)
|
|
79
|
+
ram = tier_gb * GIB
|
|
80
|
+
is_host = tier_gb == host_tier
|
|
81
|
+
if is_host:
|
|
82
|
+
recommended, source = host_recommended_bytes, "device"
|
|
83
|
+
elif device_recommended_bytes is not None:
|
|
84
|
+
recommended, source = device_recommended_bytes, "device"
|
|
85
|
+
else:
|
|
86
|
+
# Integer arithmetic (the float ratio above is for display): 2/3 up to 24 GB, 3/4 above.
|
|
87
|
+
recommended, source = (ram * 2 // 3 if tier_gb <= 24 else ram * 3 // 4), "ratio"
|
|
88
|
+
limit = min(int(1.5 * recommended), int(0.95 * ram))
|
|
89
|
+
reserve = reserve_for(tier_gb)
|
|
90
|
+
return TierLimits(
|
|
91
|
+
tier_gb=tier_gb,
|
|
92
|
+
ram_bytes=ram,
|
|
93
|
+
recommended_bytes=recommended,
|
|
94
|
+
memory_limit_bytes=limit,
|
|
95
|
+
cache_limit_bytes=min(limit, int(0.95 * recommended)),
|
|
96
|
+
wired_limit_bytes=0,
|
|
97
|
+
reserve_bytes=reserve,
|
|
98
|
+
budget_bytes=recommended,
|
|
99
|
+
ceiling_bytes=recommended - reserve,
|
|
100
|
+
budget_source=source,
|
|
101
|
+
is_host=is_host,
|
|
102
|
+
label="MEASURED" if is_host else "CAPPED",
|
|
103
|
+
)
|
|
104
|
+
|
|
105
|
+
|
|
106
|
+
def apply(limits: TierLimits) -> dict[str, int]:
|
|
107
|
+
"""Install the tier's memory, cache and wired limits; returns the previous values."""
|
|
108
|
+
return {
|
|
109
|
+
"memory": int(mx.set_memory_limit(limits.memory_limit_bytes)),
|
|
110
|
+
"cache": int(mx.set_cache_limit(limits.cache_limit_bytes)),
|
|
111
|
+
"wired": int(mx.set_wired_limit(limits.wired_limit_bytes)),
|
|
112
|
+
}
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
def current_limits() -> dict[str, int]:
|
|
116
|
+
"""The three MLX limits in force (MLX 0.32.2 has no getters; the setters return the previous value).
|
|
117
|
+
|
|
118
|
+
Each limit is set to 0 and restored to the value that returned. Verified on mlx 0.32.2: 0 is accepted
|
|
119
|
+
by all three setters, and ``set_cache_limit(0)`` does not release cached buffers on the spot (the
|
|
120
|
+
pool size was unchanged right after), so the read has no lasting side effect; call it at worker
|
|
121
|
+
start, not while a run is allocating.
|
|
122
|
+
"""
|
|
123
|
+
out: dict[str, int] = {}
|
|
124
|
+
for name, setter in (
|
|
125
|
+
("memory", mx.set_memory_limit),
|
|
126
|
+
("cache", mx.set_cache_limit),
|
|
127
|
+
("wired", mx.set_wired_limit),
|
|
128
|
+
):
|
|
129
|
+
previous = int(setter(0))
|
|
130
|
+
setter(previous)
|
|
131
|
+
out[name] = previous
|
|
132
|
+
return out
|
|
133
|
+
|
|
134
|
+
|
|
135
|
+
def limits_record(
|
|
136
|
+
limits: TierLimits,
|
|
137
|
+
*,
|
|
138
|
+
effective_memory_limit: int,
|
|
139
|
+
effective_cache_limit: int,
|
|
140
|
+
effective_wired_limit: int,
|
|
141
|
+
applied: str,
|
|
142
|
+
) -> dict[str, object]:
|
|
143
|
+
"""What a result file stores under ``limits``: the tier's numbers, the values in force, and which path installed them."""
|
|
144
|
+
return {
|
|
145
|
+
"tier": limits.as_dict(),
|
|
146
|
+
"effective": {
|
|
147
|
+
"memory_limit_bytes": effective_memory_limit,
|
|
148
|
+
"cache_limit_bytes": effective_cache_limit,
|
|
149
|
+
"wired_limit_bytes": effective_wired_limit,
|
|
150
|
+
},
|
|
151
|
+
"applied": applied,
|
|
152
|
+
}
|
|
153
|
+
|
|
154
|
+
|
|
155
|
+
__all__ = [
|
|
156
|
+
"GIB",
|
|
157
|
+
"TierLimits",
|
|
158
|
+
"apply",
|
|
159
|
+
"current_limits",
|
|
160
|
+
"host_tier_gb",
|
|
161
|
+
"limits_record",
|
|
162
|
+
"recommended_ratio",
|
|
163
|
+
"reserve_for",
|
|
164
|
+
"tier_limits",
|
|
165
|
+
]
|