From d607cc342dfa75450e02c7fb53087f8337fc1a92 Mon Sep 17 00:00:00 2001 From: Denis Ineshin Date: Fri, 18 Sep 2026 22:10:19 +0200 Subject: [PATCH 1/3] Count MLX's retained cache in the harness watchdogs The showcase, benchmark, latent-capture and A/B workers compared active memory alone against the ceiling; freed buffers sit in MLX's retained cache, resident but not active, and the default cache limit is near device memory. Both watchdogs now sample active plus cache, bound the pool at install time (2 GiB generation workers, 4 GiB decode reps so the timed decode keeps a warm pool, A/B per-condition bounds passed through), poll every 50 ms, record the policy in the abort artifact and the worker result, and abort with reason sample_error if a sample raises instead of letting the thread die silently. --- CHANGELOG.md | 15 +++ scripts/_capture_latent.py | 79 ++++++++++++--- scripts/ab_taef2_bn_domain.py | 20 ++-- scripts/bench_decode.py | 11 ++- scripts/capture_examples.py | 5 +- scripts/run_showcase.py | 73 ++++++++++++-- tests/test_ab_taef2_bn_domain.py | 41 ++++++-- tests/test_bench_decode.py | 102 +++++++++++++++++++- tests/test_capture_latent.py | 126 ++++++++++++++++++++++++ tests/test_run_showcase.py | 160 ++++++++++++++++++++++++++++++- 10 files changed, 583 insertions(+), 49 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index a11f966..641f78a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,6 +5,21 @@ All notable changes to this project will be documented in this file. The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.1.0/), and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). +## [Unreleased] + +### Changed +- The showcase, benchmark, latent-capture and A/B harness workers now count MLX's retained + buffer cache toward their memory ceiling instead of active memory alone. Buffers MLX has + freed stay resident in that cache until it is trimmed, and its default limit sits near device + memory, so a worker could hold gigabytes the watchdog never saw. Each watchdog now bounds the + cache pool when it starts (2 GiB for generation workers; 4 GiB for decode reps, which covers + their peak so the timed decode still runs from a warm pool; the A/B script keeps its own + per-condition bounds), polls every 50 ms, and records the ceiling, cache bound, cadence and + wall budget in its abort artifact and in each worker's result. A memory sample that raises + now aborts the worker with reason `sample_error` rather than ending the watchdog thread in + silence. Harness only: the library and its decode path are unchanged, and the committed + showcase numbers were not re-measured. + ## [0.8.2] - 2026-09-17 TAEF2 previews now track the full FLUX.2 VAE closely. diff --git a/scripts/_capture_latent.py b/scripts/_capture_latent.py index e33ad46..99711e6 100644 --- a/scripts/_capture_latent.py +++ b/scripts/_capture_latent.py @@ -44,6 +44,10 @@ # run_showcase.py's _MEMORY_HEADROOM_BYTES (abort at memory_size - 4 GiB). _WALL_BUDGET_S = 3600.0 _MEMORY_HEADROOM_BYTES = 4 * 1024**3 +# Same accounting as run_showcase.py's live workers: the watchdog counts MLX's retained cache +# toward the ceiling, and bounds that pool so the sum can close on a 32 GB machine. +_CAPTURE_CACHE_LIMIT_BYTES = 2 * 1024**3 +_WATCHDOG_INTERVAL_S = 0.05 def _build_argparser() -> argparse.ArgumentParser: @@ -103,10 +107,19 @@ def _abort_artifact_path(variant: str, out_dir: Path) -> Path: def _watchdog_breach_reason( - *, active_bytes: int, ceiling_bytes: int, elapsed_s: float, wall_budget_s: float + *, + active_bytes: int, + cache_bytes: int, + ceiling_bytes: int, + elapsed_s: float, + wall_budget_s: float, ) -> str | None: - """Return the first capture-run safety limit that has been breached.""" - if active_bytes >= ceiling_bytes: + """Return the first capture-run safety limit that has been breached. + + Memory is the allocator's resident footprint, active plus retained cache, as in + scripts/run_showcase.py's `_live_watchdog_breach_reason`. + """ + if active_bytes + cache_bytes >= ceiling_bytes: return "memory_ceiling" if elapsed_s > wall_budget_s: return "wall_budget" @@ -114,11 +127,18 @@ def _watchdog_breach_reason( class _CaptureWatchdog: - """Cooperatively stop a capture-run watchdog thread once generation finishes.""" + """Cooperatively stop a capture-run watchdog thread once generation finishes. + + `policy` states the limits the thread enforced (ceiling, cache bound, polling cadence, + wall budget). + """ - def __init__(self, stop_event: threading.Event, thread: threading.Thread) -> None: + def __init__( + self, stop_event: threading.Event, thread: threading.Thread, policy: dict[str, object] + ) -> None: self._stop_event = stop_event self._thread = thread + self.policy = policy def stop(self) -> None: self._stop_event.set() @@ -149,20 +169,35 @@ def _install_capture_watchdog( variant: str, out_dir: Path, *, - interval_s: float = 0.5, + interval_s: float = _WATCHDOG_INTERVAL_S, wall_budget_s: float = _WALL_BUDGET_S, + cache_limit_bytes: int = _CAPTURE_CACHE_LIMIT_BYTES, ) -> _CaptureWatchdog: """Abort a heavy capture run before it exhausts unified memory or its wall budget. - Modeled on scripts/run_showcase.py's `_install_live_watchdog`: a daemon thread polls - `mx.get_active_memory()` and elapsed wall time, and on breach writes an honest abort + Modeled on scripts/run_showcase.py's `_install_live_watchdog`: bounds MLX's retained + cache at `cache_limit_bytes`, then a daemon thread polls active plus cached allocator + bytes and elapsed wall time every `interval_s`, and on breach writes an honest abort artifact (`/.abort.json`) before killing the process with a nonzero - exit — no partial/misleading latent fixture is ever written. + exit — no partial/misleading latent fixture is ever written. A sample that raises aborts + too (reason `sample_error`) instead of leaving the run with a dead backstop. """ memory_size = int(mx.device_info().get("memory_size", 0)) if memory_size <= _MEMORY_HEADROOM_BYTES: raise RuntimeError(f"could not establish a safe memory ceiling from {memory_size} bytes") ceiling_bytes = memory_size - _MEMORY_HEADROOM_BYTES + mx.set_cache_limit(cache_limit_bytes) + policy: dict[str, object] = { + "ceiling_bytes": ceiling_bytes, + "cache_limit_bytes": cache_limit_bytes, + "interval_s": interval_s, + "wall_budget_s": wall_budget_s, + } + print( + f"watchdog[{variant}]: ceiling {ceiling_bytes / 1024**3:.1f} GiB on active+cache, " + f"cache bound {cache_limit_bytes / 1024**3:.1f} GiB, poll {interval_s} s, " + f"wall {wall_budget_s:.0f} s" + ) stop_event = threading.Event() started = time.monotonic() abort_path = _abort_artifact_path(variant, out_dir) @@ -170,9 +205,26 @@ def _install_capture_watchdog( def _watch() -> None: while not stop_event.wait(interval_s): elapsed_s = time.monotonic() - started - active_bytes = int(mx.get_active_memory()) + try: + active_bytes = int(mx.get_active_memory()) + cache_bytes = int(mx.get_cache_memory()) + except Exception as exc: + _commit_capture_watchdog_abort( + abort_path, + { + "status": "aborted", + "variant": variant, + "reason": "sample_error", + "error": f"{type(exc).__name__}: {exc}", + "elapsed_s": elapsed_s, + **policy, + }, + stop_event=stop_event, + ) + continue reason = _watchdog_breach_reason( active_bytes=active_bytes, + cache_bytes=cache_bytes, ceiling_bytes=ceiling_bytes, elapsed_s=elapsed_s, wall_budget_s=wall_budget_s, @@ -186,16 +238,17 @@ def _watch() -> None: "variant": variant, "reason": reason, "active_memory_bytes": active_bytes, - "ceiling_bytes": ceiling_bytes, + "cache_memory_bytes": cache_bytes, + "total_memory_bytes": active_bytes + cache_bytes, "elapsed_s": elapsed_s, - "wall_budget_s": wall_budget_s, + **policy, }, stop_event=stop_event, ) thread = threading.Thread(target=_watch, name=f"{variant}-capture-watchdog", daemon=True) thread.start() - return _CaptureWatchdog(stop_event, thread) + return _CaptureWatchdog(stop_event, thread, policy) def _install_memory_caps() -> None: diff --git a/scripts/ab_taef2_bn_domain.py b/scripts/ab_taef2_bn_domain.py index 8899a77..bda2a32 100644 --- a/scripts/ab_taef2_bn_domain.py +++ b/scripts/ab_taef2_bn_domain.py @@ -40,8 +40,9 @@ CONDITIONS = ("vanilla_vae", "bn_inverse", "identity") _TAEF2_CONDITIONS = ("bn_inverse", "identity") _WORKER_TIMEOUT_S = {"vanilla_vae": 1500, "bn_inverse": 300, "identity": 300} -# Every worker bounds MLX's retained-buffer pool: its default limit sits near device memory, and -# the watchdog samples active memory only. +# Every worker bounds MLX's retained-buffer pool (its default limit sits near device memory); +# the watchdog counts the pool toward its ceiling, so the full-VAE arm's 6 GiB is what keeps +# 28 GiB reachable on 32 GB while still holding that decode's transient buffers. _CACHE_LIMIT_BYTES = { "vanilla_vae": 6 * 1024**3, "bn_inverse": 2 * 1024**3, @@ -65,17 +66,17 @@ def _cache_limit_bytes(condition: str) -> int: def _install_worker_limits(condition: str) -> int: - """Pin the wired/soft memory caps and bound the MLX cache pool; return the wired cap in GB.""" - import mlx.core as mx + """Pin the wired/soft memory caps; return the wired cap in GB. + The cache bound is installed by the watchdog (`_worker_main` hands it the condition's + value), so it is recorded in the abort artifact alongside the ceiling it feeds into. + """ from scripts import bench_decode bench_condition = "vanilla_vae" if condition == "vanilla_vae" else "taef2" - installed_cap_gb = bench_decode._install_memory_caps( + return bench_decode._install_memory_caps( bench_decode._resolve_cap_gb(condition=bench_condition) ) - mx.set_cache_limit(_cache_limit_bytes(condition)) - return installed_cap_gb def _sha256(path: Path) -> str: @@ -170,7 +171,10 @@ def _worker_main(args: argparse.Namespace) -> int: import mlx.core as mx watchdog = _install_live_watchdog( - abort_path, f"ab_{condition}", wall_budget_s=_WORKER_TIMEOUT_S[condition] - 30.0 + abort_path, + f"ab_{condition}", + wall_budget_s=_WORKER_TIMEOUT_S[condition] - 30.0, + cache_limit_bytes=_cache_limit_bytes(condition), ) started = time.perf_counter() try: diff --git a/scripts/bench_decode.py b/scripts/bench_decode.py index c4bc8d7..bf13344 100644 --- a/scripts/bench_decode.py +++ b/scripts/bench_decode.py @@ -128,6 +128,13 @@ def _parse_worker_stdout(stdout: str) -> dict[str, Any]: return payload +# The decode-rep worker's cache bound. It must cover the heaviest rep's peak (the FLUX.1 VAE +# at 512x512 reaches ~3.7 GiB) so nothing is evicted between the warmup and the timed decode, +# or the committed steady-state timings would move; 28 GiB ceiling - 4 GiB pool still leaves +# every rep far below the memory arm. +_BENCH_CACHE_LIMIT_BYTES = 4 * 1024**3 + + def _watchdog_abort_path(save_to: Path) -> Path: """Where this rep's watchdog abort artifact would be written, if any. @@ -457,7 +464,7 @@ def _save_webp(image_uint8_nhwc: Any, target: Path) -> None: def _worker_main(args: argparse.Namespace) -> int: """Run one (condition, rep) inside this subprocess. Emit sentinel. - Installs the same active-memory watchdog the live-scenario worker path uses + Installs the same active-plus-cache memory watchdog the live-scenario worker path uses (scripts.run_showcase._install_live_watchdog) around model construction + decode, so a rep that would otherwise page-storm the machine aborts with an honest artifact (see _watchdog_abort_path) and a nonzero exit instead of risking a kernel panic. @@ -476,6 +483,7 @@ def _worker_main(args: argparse.Namespace) -> int: _watchdog_abort_path(args.save_to), f"{args.condition}_rep{args.rep}", wall_budget_s=_worker_wall_budget_s(args.condition), + cache_limit_bytes=_BENCH_CACHE_LIMIT_BYTES, ) try: import mlx.core as mx @@ -521,6 +529,7 @@ def _worker_main(args: argparse.Namespace) -> int: "image_path": _repo_relative(args.save_to), "requested_cap_gb": args.applied_cap_gb, "installed_cap_gb": installed_cap_gb, + "watchdog": watchdog.policy, } ) ) diff --git a/scripts/capture_examples.py b/scripts/capture_examples.py index 3fb0ad2..4a9f43c 100644 --- a/scripts/capture_examples.py +++ b/scripts/capture_examples.py @@ -89,7 +89,8 @@ watchdog (`scripts._capture_latent._install_capture_watchdog`, imported directly rather than re-implemented — same discipline as `scripts/run_showcase.py`'s `_install_live_watchdog`) aborts the run — writing `/.abort.json` and exiting nonzero via -`os._exit(70)` — if active memory nears the device ceiling or the wall budget is exceeded. A +`os._exit(70)` — if MLX's active plus retained-cache memory nears the device ceiling or the +wall budget is exceeded (the watchdog also bounds that cache pool). A stale abort artifact from a prior aborted run is cleared before each new attempt. Publishing is atomic (mirrors the converted-weights cache's temp-then-rename discipline, see @@ -134,7 +135,7 @@ module names `img_in`/`txt_in`/`time_text_embed`/`proj_out`/`norm_out`). Pass `--qwen-uniform-q4` for the plain uniform-q4 build instead. The paper reports about 1.9 GiB of extra peak MLX allocation for the mixed build (27.6 -> 29.5 GiB at 512x512, 50-step CFG); this -script's 20-step, no-CFG run should sit lower, but if the active-memory watchdog aborts a +script's 20-step, no-CFG run should sit lower, but if the memory watchdog aborts a mixed-precision capture, that is an honest signal to report, not a reason to raise the memory ceiling. diff --git a/scripts/run_showcase.py b/scripts/run_showcase.py index ace21b0..a9eab9d 100644 --- a/scripts/run_showcase.py +++ b/scripts/run_showcase.py @@ -445,6 +445,13 @@ def _validate_live_artifacts( _LIVE_SCENARIOS = frozenset({"live_preview", "zimage_live_preview", "combined"}) _LIVE_WALL_BUDGET_S = 3300.0 _MEMORY_HEADROOM_BYTES = 4 * 1024**3 +# Every model-loading worker bounds MLX's retained-buffer pool. Freed buffers sit in that pool, +# resident but not "active", and its default limit is near device memory; the watchdog counts +# the pool toward the ceiling, so the bound is what makes the arithmetic close on 32 GB: +# the heaviest live scenario (Z-Image) peaks near 25.9 GiB active, and 25.9 + 2 stays under +# the 28 GiB ceiling. The decode-rep worker uses a larger bound (see scripts/bench_decode.py). +_LIVE_CACHE_LIMIT_BYTES = 2 * 1024**3 +_WATCHDOG_INTERVAL_S = 0.05 def _all_scenario_order() -> list[str]: @@ -638,12 +645,18 @@ def _write_report(path: Path, report: dict[str, Any]) -> None: def _live_watchdog_breach_reason( *, active_bytes: int, + cache_bytes: int, ceiling_bytes: int, elapsed_s: float, wall_budget_s: float, ) -> str | None: - """Return the first live-worker safety limit that has been breached.""" - if active_bytes >= ceiling_bytes: + """Return the first live-worker safety limit that has been breached. + + The memory arm compares the MLX allocator's resident footprint, active plus retained + cache, against the ceiling: cached buffers are freed from the graph's point of view but + still occupy unified memory until MLX releases them (`mx.clear_cache` or the cache limit). + """ + if active_bytes + cache_bytes >= ceiling_bytes: return "memory_ceiling" if elapsed_s > wall_budget_s: return "wall_budget" @@ -680,11 +693,18 @@ def _resolve_override_wired_gb(cap_gb: int, device_wired_gb: int) -> int | None: class _LiveWatchdog: - """Cooperatively stop a live-worker watchdog thread after generation.""" + """Cooperatively stop a live-worker watchdog thread after generation. - def __init__(self, stop_event: threading.Event, thread: threading.Thread) -> None: + `policy` states the limits the thread enforced (ceiling, cache bound, polling cadence, + wall budget) so a worker can record them next to its result. + """ + + def __init__( + self, stop_event: threading.Event, thread: threading.Thread, policy: dict[str, Any] + ) -> None: self._stop_event = stop_event self._thread = thread + self.policy = policy def stop(self) -> None: self._stop_event.set() @@ -695,25 +715,56 @@ def _install_live_watchdog( result_path: Path, scenario: str, *, - interval_s: float = 0.5, + interval_s: float = _WATCHDOG_INTERVAL_S, wall_budget_s: float = _LIVE_WALL_BUDGET_S, + cache_limit_bytes: int = _LIVE_CACHE_LIMIT_BYTES, ) -> _LiveWatchdog: - """Abort a live worker before it exhausts unified memory or its wall budget.""" + """Abort a live worker before it exhausts unified memory or its wall budget. + + Bounds MLX's retained-buffer pool at `cache_limit_bytes` first, then polls active plus + cached allocator bytes every `interval_s` (the poll runs while `mx.eval` holds no GIL, so + a fast cadence is cheap). A sample that raises is itself a breach: the worker aborts + with reason `sample_error` rather than running on with a dead backstop. + """ import mlx.core as mx memory_size = int(mx.device_info().get("memory_size", 0)) if memory_size <= _MEMORY_HEADROOM_BYTES: raise TaefError(f"could not establish a safe memory ceiling from {memory_size} bytes") ceiling_bytes = memory_size - _MEMORY_HEADROOM_BYTES + mx.set_cache_limit(cache_limit_bytes) + policy: dict[str, Any] = { + "ceiling_bytes": ceiling_bytes, + "cache_limit_bytes": cache_limit_bytes, + "interval_s": interval_s, + "wall_budget_s": wall_budget_s, + } stop_event = threading.Event() started = time.monotonic() def _watch() -> None: while not stop_event.wait(interval_s): elapsed_s = time.monotonic() - started - active_bytes = int(mx.get_active_memory()) + try: + active_bytes = int(mx.get_active_memory()) + cache_bytes = int(mx.get_cache_memory()) + except Exception as exc: + _commit_watchdog_abort( + result_path, + { + "status": "aborted", + "scenario": scenario, + "reason": "sample_error", + "error": f"{type(exc).__name__}: {exc}", + "elapsed_s": elapsed_s, + **policy, + }, + stop_event=stop_event, + ) + continue reason = _live_watchdog_breach_reason( active_bytes=active_bytes, + cache_bytes=cache_bytes, ceiling_bytes=ceiling_bytes, elapsed_s=elapsed_s, wall_budget_s=wall_budget_s, @@ -727,16 +778,17 @@ def _watch() -> None: "scenario": scenario, "reason": reason, "active_memory_bytes": active_bytes, - "ceiling_bytes": ceiling_bytes, + "cache_memory_bytes": cache_bytes, + "total_memory_bytes": active_bytes + cache_bytes, "elapsed_s": elapsed_s, - "wall_budget_s": wall_budget_s, + **policy, }, stop_event=stop_event, ) thread = threading.Thread(target=_watch, name=f"{scenario}-watchdog", daemon=True) thread.start() - return _LiveWatchdog(stop_event, thread) + return _LiveWatchdog(stop_event, thread, policy) # --------------------------------------------------------------------------- @@ -1043,6 +1095,7 @@ def main(argv: list[str] | None = None) -> int: result = _SCENARIO_DISPATCH[args.live_worker](args) finally: watchdog.stop() + result["watchdog"] = watchdog.policy _write_report(args.live_result, result) return 0 diff --git a/tests/test_ab_taef2_bn_domain.py b/tests/test_ab_taef2_bn_domain.py index d8f03e8..18c7bbf 100644 --- a/tests/test_ab_taef2_bn_domain.py +++ b/tests/test_ab_taef2_bn_domain.py @@ -93,21 +93,42 @@ def test_pending_units_reruns_results_that_belong_to_a_different_latent(tmp_path assert ab._pending_units(tmp_path, "lat", latent_sha256="bbbb") == list(ab.CONDITIONS) -def test_every_worker_installs_a_cache_limit_before_loading_anything(monkeypatch) -> None: +def test_every_worker_hands_its_cache_bound_to_the_watchdog(monkeypatch, tmp_path) -> None: """Catches: a worker (the full-VAE arm is the heaviest) running with MLX's near-device-size - default cache limit, where retained buffers sit far above the active-memory watchdog's view.""" + default cache limit, or with the watchdog's generic bound instead of the condition's own + (6 GiB for the full VAE, 2 GiB for the TAEF2 arms). The watchdog installs the bound, so + the worker must pass it; a second bare `mx.set_cache_limit` call would be overridden.""" + import argparse + import scripts.ab_taef2_bn_domain as ab - import scripts.bench_decode as bench + import scripts.run_showcase as rs + + class _StopError(Exception): + pass + + class _FakeWatchdog: + def stop(self) -> None: + pass - calls: list[tuple[str, int]] = [] - monkeypatch.setattr(bench, "_install_memory_caps", lambda cap: calls.append(("caps", cap)) or 7) - monkeypatch.setattr(mx, "set_cache_limit", lambda n: calls.append(("cache", n))) + install_calls: list[dict[str, object]] = [] + + def _fake_install(result_path, scenario: str, **kwargs: object) -> _FakeWatchdog: + install_calls.append(kwargs) + return _FakeWatchdog() + + monkeypatch.setattr(rs, "_install_live_watchdog", _fake_install) + monkeypatch.setattr(ab, "_install_worker_limits", lambda condition: 7) + monkeypatch.setattr(mx, "load", lambda path: (_ for _ in ()).throw(_StopError())) for condition in ab.CONDITIONS: - calls.clear() - assert ab._install_worker_limits(condition) == 7 - assert [name for name, _ in calls] == ["caps", "cache"] - assert 0 < calls[1][1] <= 8 * 1024**3 + install_calls.clear() + args = argparse.Namespace( + worker=condition, latent=tmp_path / "lat.safetensors", out_dir=tmp_path + ) + with pytest.raises(_StopError): + ab._worker_main(args) + assert install_calls[0]["cache_limit_bytes"] == ab._cache_limit_bytes(condition) + assert 0 < ab._cache_limit_bytes(condition) <= 8 * 1024**3 def test_score_keeps_a_computed_lpips_when_the_other_arm_fails(tmp_path: Path, monkeypatch) -> None: diff --git a/tests/test_bench_decode.py b/tests/test_bench_decode.py index 7366d32..4d88a8e 100644 --- a/tests/test_bench_decode.py +++ b/tests/test_bench_decode.py @@ -6,6 +6,7 @@ import json from pathlib import Path +from typing import ClassVar from unittest.mock import patch import pytest @@ -348,6 +349,8 @@ def test_worker_main_clears_stale_watchdog_abort_artifact_before_running( ) class _FakeWatchdog: + policy: ClassVar[dict[str, object]] = {"cache_limit_bytes": 0, "interval_s": 0.05} + def stop(self) -> None: pass @@ -397,6 +400,8 @@ def _fake_decode() -> str: return "IMG" class _FakeWatchdog: + policy: ClassVar[dict[str, object]] = {"cache_limit_bytes": 0, "interval_s": 0.05} + def stop(self) -> None: pass @@ -463,6 +468,8 @@ def test_worker_main_installs_watchdog_with_condition_scoped_wall_budget( ) class _FakeWatchdog: + policy: ClassVar[dict[str, object]] = {"cache_limit_bytes": 0, "interval_s": 0.05} + def stop(self) -> None: pass @@ -488,7 +495,7 @@ def _fake_install(result_path: Path, scenario: str, **kwargs: object) -> _FakeWa ) assert bench._worker_main(args) == 0 - assert install_calls == [{"wall_budget_s": bench._worker_wall_budget_s("taef1")}] + assert [c["wall_budget_s"] for c in install_calls] == [bench._worker_wall_budget_s("taef1")] def test_prep_taef2_decodes_the_normalized_latent(monkeypatch) -> None: @@ -518,3 +525,96 @@ def decode_image(self, nhwc: mx.array) -> mx.array: expected = unpack_flux2_latent(packed, latent_height=2, latent_width=3) assert np.array_equal(np.array(seen["nhwc"]), np.array(expected)) + + +def test_worker_main_passes_a_measurement_preserving_cache_bound_to_the_watchdog( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """Catches: the decode-rep worker inheriting the live workers' 2 GiB cache bound. The + full-VAE reps peak at ~3.7 GiB, so a pool smaller than that evicts buffers between the + warmup and the timed decode and shifts the committed steady-state timings; the bench + bound must cover that peak while still leaving the 28 GiB ceiling reachable.""" + import argparse + + import mlx.core as mx + import scripts.bench_decode as bench + + latent_file = tmp_path / "latent.safetensors" + mx.save_safetensors( + str(latent_file), + {"latent": mx.zeros((1, 2, 2, 16)), "height": mx.array(16), "width": mx.array(16)}, + ) + + class _FakeWatchdog: + policy: ClassVar[dict[str, object]] = {"cache_limit_bytes": 4 * 1024**3, "interval_s": 0.05} + + def stop(self) -> None: + pass + + install_calls: list[dict[str, object]] = [] + + def _fake_install(result_path: Path, scenario: str, **kwargs: object) -> _FakeWatchdog: + install_calls.append(kwargs) + return _FakeWatchdog() + + monkeypatch.setattr("scripts.run_showcase._install_live_watchdog", _fake_install) + monkeypatch.setattr(bench, "_install_memory_caps", lambda cap: 1) + monkeypatch.setattr(bench, "_prep_taef1", lambda latent, h, w: lambda: "IMG") + monkeypatch.setattr(bench, "_measure_steady_state", lambda decode_fn: ("IMG", 0.25, 2.0)) + monkeypatch.setattr(bench, "_save_webp", lambda image, target: None) + + args = argparse.Namespace( + condition="taef1", + rep=0, + latent=latent_file, + save_to=tmp_path / "out.webp", + applied_cap_gb=1, + flux_variant="flux1-dev", + ) + + assert bench._worker_main(args) == 0 + assert install_calls[0]["cache_limit_bytes"] == bench._BENCH_CACHE_LIMIT_BYTES + assert 4 * 1024**3 <= bench._BENCH_CACHE_LIMIT_BYTES <= 8 * 1024**3 + + +def test_worker_main_records_the_watchdog_policy_in_the_sentinel( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str] +) -> None: + """Catches: a rep result that states the wired cap it ran under but not the cache bound + or the polling cadence, so the committed report cannot reproduce the run's policy.""" + import argparse + + import mlx.core as mx + import scripts.bench_decode as bench + + latent_file = tmp_path / "latent.safetensors" + mx.save_safetensors( + str(latent_file), + {"latent": mx.zeros((1, 2, 2, 16)), "height": mx.array(16), "width": mx.array(16)}, + ) + + class _FakeWatchdog: + policy: ClassVar[dict[str, object]] = {"cache_limit_bytes": 4 * 1024**3, "interval_s": 0.05} + + def stop(self) -> None: + pass + + monkeypatch.setattr( + "scripts.run_showcase._install_live_watchdog", lambda *a, **kw: _FakeWatchdog() + ) + monkeypatch.setattr(bench, "_install_memory_caps", lambda cap: 1) + monkeypatch.setattr(bench, "_prep_taef1", lambda latent, h, w: lambda: "IMG") + monkeypatch.setattr(bench, "_measure_steady_state", lambda decode_fn: ("IMG", 0.25, 2.0)) + monkeypatch.setattr(bench, "_save_webp", lambda image, target: None) + + args = argparse.Namespace( + condition="taef1", + rep=0, + latent=latent_file, + save_to=tmp_path / "out.webp", + applied_cap_gb=1, + flux_variant="flux1-dev", + ) + assert bench._worker_main(args) == 0 + payload = bench._parse_worker_stdout(capsys.readouterr().out) + assert payload["watchdog"] == {"cache_limit_bytes": 4 * 1024**3, "interval_s": 0.05} diff --git a/tests/test_capture_latent.py b/tests/test_capture_latent.py index fac92f5..7749b90 100644 --- a/tests/test_capture_latent.py +++ b/tests/test_capture_latent.py @@ -235,3 +235,129 @@ def _broken_write(self: object, text: str) -> None: tmp_path / "r.abort.json", {"status": "aborted"}, stop_event=threading.Event() ) assert exits == [70] + + +def test_capture_watchdog_breach_reason_counts_retained_cache_toward_the_ceiling() -> None: + """Catches: the capture ceiling compared against active memory alone (see the same + test for scripts/run_showcase.py; the two watchdogs must agree on the accounting).""" + import scripts._capture_latent as cl + + assert ( + cl._watchdog_breach_reason( + active_bytes=20, cache_bytes=8, ceiling_bytes=28, elapsed_s=0, wall_budget_s=5 + ) + == "memory_ceiling" + ) + assert ( + cl._watchdog_breach_reason( + active_bytes=20, cache_bytes=7, ceiling_bytes=28, elapsed_s=0, wall_budget_s=5 + ) + is None + ) + assert ( + cl._watchdog_breach_reason( + active_bytes=20, cache_bytes=7, ceiling_bytes=28, elapsed_s=6, wall_budget_s=5 + ) + == "wall_budget" + ) + + +def _install_capture_watchdog_with_fake_mlx( + monkeypatch, tmp_path: Path, *, active_bytes: object, cache_bytes: int +) -> tuple[object, list[dict[str, object]]]: + """Run the real capture watchdog thread against a fake MLX memory API (32 GiB device). + + `active_bytes` may be an exception instance, raised by the active-memory sample. The + abort commit is replaced by a recorder that stops the thread, so exactly one payload is + observed and the process is never exited. + """ + import threading + + import scripts._capture_latent as cl + + monkeypatch.setattr(cl.mx, "device_info", lambda: {"memory_size": 32 * 1024**3}) + monkeypatch.setattr(cl.mx, "set_cache_limit", lambda n: 0) + + def _active() -> int: + if isinstance(active_bytes, BaseException): + raise active_bytes + return int(active_bytes) # type: ignore[call-overload] + + monkeypatch.setattr(cl.mx, "get_active_memory", _active) + monkeypatch.setattr(cl.mx, "get_cache_memory", lambda: cache_bytes) + + payloads: list[dict[str, object]] = [] + + def _record_abort( + abort_path: Path, payload: dict[str, object], *, stop_event: threading.Event + ) -> None: + payloads.append(payload) + stop_event.set() + + monkeypatch.setattr(cl, "_commit_capture_watchdog_abort", _record_abort) + watchdog = cl._install_capture_watchdog("flux1-dev", tmp_path, interval_s=0.005) + watchdog._thread.join(timeout=5) + assert not watchdog._thread.is_alive(), "watchdog thread never reached the abort path" + return watchdog, payloads + + +def test_capture_watchdog_aborts_on_active_plus_cache_and_records_both( + monkeypatch, tmp_path: Path +) -> None: + """Catches: `_watch` sampling `mx.get_active_memory()` only — 20 GiB active under a + 28 GiB ceiling, but 29 GiB with the retained cache counted.""" + active = 20 * 1024**3 + cache = 9 * 1024**3 + _, payloads = _install_capture_watchdog_with_fake_mlx( + monkeypatch, tmp_path, active_bytes=active, cache_bytes=cache + ) + + assert len(payloads) == 1 + payload = payloads[0] + assert payload["reason"] == "memory_ceiling" + assert payload["variant"] == "flux1-dev" + assert payload["active_memory_bytes"] == active + assert payload["cache_memory_bytes"] == cache + assert payload["total_memory_bytes"] == active + cache + assert payload["ceiling_bytes"] == 28 * 1024**3 + + +def test_capture_watchdog_aborts_explicitly_when_a_memory_sample_fails( + monkeypatch, tmp_path: Path +) -> None: + """Catches: a sampling exception killing the daemon thread silently and leaving an + hour-long capture with no memory backstop.""" + _, payloads = _install_capture_watchdog_with_fake_mlx( + monkeypatch, tmp_path, active_bytes=RuntimeError("metal device gone"), cache_bytes=0 + ) + + assert len(payloads) == 1 + assert payloads[0]["reason"] == "sample_error" + assert "metal device gone" in str(payloads[0]["error"]) + assert "active_memory_bytes" not in payloads[0] + + +def test_capture_watchdog_bounds_the_cache_pool_and_records_its_policy( + monkeypatch, tmp_path: Path +) -> None: + """Catches: a capture run under MLX's default (near device-size) cache limit, and a + polling cadence that regressed to the old 0.5 s.""" + import scripts._capture_latent as cl + + limits: list[int] = [] + monkeypatch.setattr(cl.mx, "device_info", lambda: {"memory_size": 32 * 1024**3}) + monkeypatch.setattr(cl.mx, "set_cache_limit", lambda n: limits.append(n) or 0) + monkeypatch.setattr(cl.mx, "get_active_memory", lambda: 0) + monkeypatch.setattr(cl.mx, "get_cache_memory", lambda: 0) + + watchdog = cl._install_capture_watchdog("flux1-dev", tmp_path) + watchdog.stop() + + assert limits == [cl._CAPTURE_CACHE_LIMIT_BYTES] + assert 0 < cl._CAPTURE_CACHE_LIMIT_BYTES <= 2 * 1024**3 + assert watchdog.policy == { + "ceiling_bytes": 28 * 1024**3, + "cache_limit_bytes": cl._CAPTURE_CACHE_LIMIT_BYTES, + "interval_s": 0.05, + "wall_budget_s": cl._WALL_BUDGET_S, + } diff --git a/tests/test_run_showcase.py b/tests/test_run_showcase.py index 53573c7..2305706 100644 --- a/tests/test_run_showcase.py +++ b/tests/test_run_showcase.py @@ -2,6 +2,7 @@ import json from pathlib import Path +from typing import ClassVar import pytest @@ -104,6 +105,8 @@ def test_live_worker_mode_runs_one_raw_scenario(tmp_path: Path, monkeypatch) -> watchdog_events: list[str] = [] class _FakeWatchdog: + policy: ClassVar[dict[str, object]] = {"cache_limit_bytes": 1, "interval_s": 0.05} + def stop(self) -> None: watchdog_events.append("stopped") @@ -134,7 +137,12 @@ def _fake_install(result_path: Path, scenario: str) -> _FakeWatchdog: ) assert calls == ["live_preview"] assert watchdog_events == ["installed:live_preview", "stopped"] - assert json.loads(result_path.read_text()) == {"status": "ok"} + # The worker result carries the watchdog policy it ran under, so the committed report + # states the cache bound and polling cadence next to the wired cap. + assert json.loads(result_path.read_text()) == { + "status": "ok", + "watchdog": {"cache_limit_bytes": 1, "interval_s": 0.05}, + } def test_hardware_metadata_names_generation_and_decode_dtypes() -> None: @@ -168,24 +176,166 @@ def test_live_watchdog_breach_reason_checks_memory_before_wall() -> None: assert ( _live_watchdog_breach_reason( - active_bytes=28, ceiling_bytes=28, elapsed_s=10, wall_budget_s=5 + active_bytes=28, cache_bytes=0, ceiling_bytes=28, elapsed_s=10, wall_budget_s=5 ) == "memory_ceiling" ) assert ( _live_watchdog_breach_reason( - active_bytes=20, ceiling_bytes=28, elapsed_s=6, wall_budget_s=5 + active_bytes=20, cache_bytes=0, ceiling_bytes=28, elapsed_s=6, wall_budget_s=5 ) == "wall_budget" ) assert ( _live_watchdog_breach_reason( - active_bytes=20, ceiling_bytes=28, elapsed_s=4, wall_budget_s=5 + active_bytes=20, cache_bytes=0, ceiling_bytes=28, elapsed_s=4, wall_budget_s=5 + ) + is None + ) + + +def test_live_watchdog_breach_reason_counts_retained_cache_toward_the_ceiling() -> None: + """Catches: the ceiling compared against active memory alone. MLX keeps freed buffers in + a retained pool that is resident but not "active"; 20 GiB active + 8 GiB cache is a + 28 GiB footprint and must trip a 28 GiB ceiling.""" + from scripts.run_showcase import _live_watchdog_breach_reason + + assert ( + _live_watchdog_breach_reason( + active_bytes=20, cache_bytes=8, ceiling_bytes=28, elapsed_s=0, wall_budget_s=5 + ) + == "memory_ceiling" + ) + assert ( + _live_watchdog_breach_reason( + active_bytes=20, cache_bytes=7, ceiling_bytes=28, elapsed_s=0, wall_budget_s=5 ) is None ) +def _install_live_watchdog_with_fake_mlx( + monkeypatch, + tmp_path: Path, + *, + active_bytes: object, + cache_bytes: int, + **install_kwargs: object, +) -> tuple[object, list[dict[str, object]]]: + """Run the real watchdog thread against a fake MLX memory API (memory_size 32 GiB). + + `active_bytes` may be an exception instance, in which case the active-memory sample + raises it. The abort commit is replaced by a recorder that also stops the thread, so + the test observes exactly one payload and the process is never exited. + """ + import threading + + import mlx.core as mx + import scripts.run_showcase as rs + + monkeypatch.setattr(mx, "device_info", lambda: {"memory_size": 32 * 1024**3}) + monkeypatch.setattr(mx, "set_cache_limit", lambda n: 0) + + def _active() -> int: + if isinstance(active_bytes, BaseException): + raise active_bytes + return int(active_bytes) # type: ignore[call-overload] + + monkeypatch.setattr(mx, "get_active_memory", _active) + monkeypatch.setattr(mx, "get_cache_memory", lambda: cache_bytes) + + payloads: list[dict[str, object]] = [] + + def _record_abort( + result_path: Path, payload: dict[str, object], *, stop_event: threading.Event + ) -> None: + payloads.append(payload) + stop_event.set() + + monkeypatch.setattr(rs, "_commit_watchdog_abort", _record_abort) + watchdog = rs._install_live_watchdog( + tmp_path / "r.json", "scn", interval_s=0.005, **install_kwargs + ) + watchdog._thread.join(timeout=5) + assert not watchdog._thread.is_alive(), "watchdog thread never reached the abort path" + return watchdog, payloads + + +def test_live_watchdog_aborts_on_active_plus_cache_and_records_both( + monkeypatch, tmp_path: Path +) -> None: + """Catches: `_watch` sampling `mx.get_active_memory()` only. 20 GiB active is under a + 28 GiB ceiling; with 9 GiB retained cache the footprint is 29 GiB and the worker must + abort with the two components and their sum recorded separately.""" + active = 20 * 1024**3 + cache = 9 * 1024**3 + _, payloads = _install_live_watchdog_with_fake_mlx( + monkeypatch, tmp_path, active_bytes=active, cache_bytes=cache + ) + + assert len(payloads) == 1 + payload = payloads[0] + assert payload["reason"] == "memory_ceiling" + assert payload["active_memory_bytes"] == active + assert payload["cache_memory_bytes"] == cache + assert payload["total_memory_bytes"] == active + cache + assert payload["ceiling_bytes"] == 28 * 1024**3 + + +def test_live_watchdog_aborts_explicitly_when_a_memory_sample_fails( + monkeypatch, tmp_path: Path +) -> None: + """Catches: an exception in the sampling call killing the daemon thread silently, which + leaves the worker running with no memory backstop at all. A failed observation must + be an explicit abort that names the error, never a zero sample and never silence.""" + _, payloads = _install_live_watchdog_with_fake_mlx( + monkeypatch, + tmp_path, + active_bytes=RuntimeError("metal device gone"), + cache_bytes=0, + ) + + assert len(payloads) == 1 + assert payloads[0]["reason"] == "sample_error" + assert "metal device gone" in str(payloads[0]["error"]) + assert "active_memory_bytes" not in payloads[0] + + +def test_live_watchdog_bounds_the_cache_pool_before_polling_and_records_its_policy( + monkeypatch, tmp_path: Path +) -> None: + """Catches: a live worker running with MLX's near-device-size default cache limit (the + accounting can then never close on 32 GB), and a policy the report cannot reproduce. + The bound is installed by the watchdog so every model-loading harness path gets it.""" + import mlx.core as mx + import scripts.run_showcase as rs + + limits: list[int] = [] + monkeypatch.setattr(mx, "device_info", lambda: {"memory_size": 32 * 1024**3}) + monkeypatch.setattr(mx, "set_cache_limit", lambda n: limits.append(n) or 0) + monkeypatch.setattr(mx, "get_active_memory", lambda: 0) + monkeypatch.setattr(mx, "get_cache_memory", lambda: 0) + + watchdog = rs._install_live_watchdog(tmp_path / "r.json", "scn", cache_limit_bytes=3 * 1024**3) + watchdog.stop() + + assert limits == [3 * 1024**3] + assert watchdog.policy == { + "ceiling_bytes": 28 * 1024**3, + "cache_limit_bytes": 3 * 1024**3, + "interval_s": 0.05, + "wall_budget_s": rs._LIVE_WALL_BUDGET_S, + } + + +def test_live_watchdog_default_cache_bound_closes_the_ceiling_arithmetic() -> None: + """Catches: a default cache bound so large that a live worker's measured peak plus a + full cache pool overshoots the 28 GiB ceiling on 32 GB (Z-Image peaks at ~25.9 GiB).""" + import scripts.run_showcase as rs + + assert 0 < rs._LIVE_CACHE_LIMIT_BYTES <= 2 * 1024**3 + + def test_json_schema_rejects_unknown_version(tmp_path: Path) -> None: from scripts.run_showcase import _load_report @@ -607,6 +757,8 @@ def test_vs_vae_worker_installs_active_memory_watchdog( watchdog_events: list[str] = [] class _FakeWatchdog: + policy: ClassVar[dict[str, object]] = {"cache_limit_bytes": 0, "interval_s": 0.05} + def stop(self) -> None: watchdog_events.append("stopped") From f1387ab520b082e98e35c5e689f3b60b7dfbe08a Mon Sep 17 00:00:00 2001 From: Denis Ineshin Date: Fri, 18 Sep 2026 22:31:03 +0200 Subject: [PATCH 2/3] Record what the harness watchdogs observe and say what the bound is for The cache bound does not close arithmetic the memory limit already closes (mlx 0.32.2 drains the retained cache before an allocation that would cross it), so the comments, CHANGELOG and COMPARISON now say what the change is: an honest active-plus-cache reading with a stated policy. Each watchdog keeps high-water marks of the sum it compares and of the cache term, recorded as watchdog_observed in the live result, the bench sentinel and its report block, and the A/B unit result. The orchestrators print every term of an abort plus the error text of a failed sample. The live bound rises to 4 GiB after an interleaved timing check of live_preview at 4 and 20 GiB (10.76 s vs 11.07 s median, three reps each); the capture bound follows it. The A/B condition-to-cap mapping is tested for real again. --- CHANGELOG.md | 21 ++++--- COMPARISON.md | 2 +- scripts/_capture_latent.py | 38 ++++++++--- scripts/ab_taef2_bn_domain.py | 9 ++- scripts/bench_decode.py | 30 ++++++--- scripts/run_showcase.py | 71 ++++++++++++++++----- tests/test_ab_taef2_bn_domain.py | 55 ++++++++++++++++ tests/test_bench_decode.py | 95 +++++++++++++++++++++++++++- tests/test_capture_latent.py | 31 ++++++++- tests/test_run_showcase.py | 105 +++++++++++++++++++++++++++---- 10 files changed, 397 insertions(+), 60 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 641f78a..69ec95f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,16 +9,17 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Changed - The showcase, benchmark, latent-capture and A/B harness workers now count MLX's retained - buffer cache toward their memory ceiling instead of active memory alone. Buffers MLX has - freed stay resident in that cache until it is trimmed, and its default limit sits near device - memory, so a worker could hold gigabytes the watchdog never saw. Each watchdog now bounds the - cache pool when it starts (2 GiB for generation workers; 4 GiB for decode reps, which covers - their peak so the timed decode still runs from a warm pool; the A/B script keeps its own - per-condition bounds), polls every 50 ms, and records the ceiling, cache bound, cadence and - wall budget in its abort artifact and in each worker's result. A memory sample that raises - now aborts the worker with reason `sample_error` rather than ending the watchdog thread in - silence. Harness only: the library and its decode path are unchanged, and the committed - showcase numbers were not re-measured. + buffer cache toward their memory ceiling instead of active memory alone, which is the number + MLX itself reports as resident. Each watchdog bounds the cache pool when it starts and states + that bound, its polling cadence (now 50 ms), the ceiling and the wall budget in its abort + artifact and in each worker's result, together with the highest active-plus-cache reading it + saw, so a report row can say how close a run came to the ceiling rather than only its active + peak. The bounds are 4 GiB for generation workers and decode reps (the reps peak at 3.7 GiB, + so the timed decode still runs from a warm pool) and the A/B script's existing per-condition + values. A memory sample that raises now aborts the worker with reason `sample_error` and the + error text rather than ending the watchdog thread in silence, and the orchestrators print + every term the watchdog compared. Harness only: the library and its decode path are + unchanged, and the committed showcase numbers were not re-measured. ## [0.8.2] - 2026-09-17 diff --git a/COMPARISON.md b/COMPARISON.md index 60d614e..4afab01 100644 --- a/COMPARISON.md +++ b/COMPARISON.md @@ -10,7 +10,7 @@ Visual showcase of what mlx-taef does on real generations. Every number on this - mlx-taef source `v0.7.1-8-g28af6c5` at commit `28af6c5`; installed distribution `0.7.2.dev4+ga4e5df5eb.d20260809` - The three FLUX.2 scenarios (`taef2_vs_vae`, `live_preview`, `combined`) were re-measured on 2026-09-17 under macOS Darwin 27.0.0, mflux 0.19.1, MLX 0.32.2 and mlx-taef `v0.8.1-5-g5d97cd5`, after the TAEF2 input fix described below. The report records that run under `scenario_updates`; the FLUX.1 and Z-Image scenarios keep the run described above. - Quantization: int4 (mflux `quantize=4`), bf16 generation, fp32 tiny-autoencoder decode -- Every condition ran in an isolated subprocess with `mx.set_wired_limit` set per the cap column. Each decode was timed after one untimed warmup call, so the figure reflects steady-state per-step decode rather than a cold first call. Every model-loading subprocess — the three live-generation workers and the vs-VAE decode reps alike — enforces a 28 GiB active-memory ceiling; the live-generation workers add a 55-minute wall budget on top. Hardware metadata is recorded inline in `_artifacts/showcase_report.json`. +- Every condition ran in an isolated subprocess with `mx.set_wired_limit` set per the cap column. Each decode was timed after one untimed warmup call, so the figure reflects steady-state per-step decode rather than a cold first call. Every model-loading subprocess — the three live-generation workers and the vs-VAE decode reps alike — enforces a 28 GiB ceiling on MLX active memory (runs after 2026-09-18 count active plus retained cache against the same ceiling and record the reading; the numbers here predate that change); the live-generation workers add a 55-minute wall budget on top. Hardware metadata is recorded inline in `_artifacts/showcase_report.json`. ## Where this fits diff --git a/scripts/_capture_latent.py b/scripts/_capture_latent.py index 99711e6..ac112c3 100644 --- a/scripts/_capture_latent.py +++ b/scripts/_capture_latent.py @@ -45,8 +45,10 @@ _WALL_BUDGET_S = 3600.0 _MEMORY_HEADROOM_BYTES = 4 * 1024**3 # Same accounting as run_showcase.py's live workers: the watchdog counts MLX's retained cache -# toward the ceiling, and bounds that pool so the sum can close on a 32 GB machine. -_CAPTURE_CACHE_LIMIT_BYTES = 2 * 1024**3 +# toward the ceiling and states the bound it installed; the memory limit already drains the +# pool under pressure, so the bound matters in the low-active phases and for the record. It +# tracks run_showcase._LIVE_CACHE_LIMIT_BYTES (captures run the same generation recipes). +_CAPTURE_CACHE_LIMIT_BYTES = 4 * 1024**3 _WATCHDOG_INTERVAL_S = 0.05 @@ -116,8 +118,9 @@ def _watchdog_breach_reason( ) -> str | None: """Return the first capture-run safety limit that has been breached. - Memory is the allocator's resident footprint, active plus retained cache, as in - scripts/run_showcase.py's `_live_watchdog_breach_reason`. + Memory is the MLX allocator's resident footprint, active plus retained cache, as in + scripts/run_showcase.py's `_live_watchdog_breach_reason`; non-MLX process memory is + covered only by the 4 GiB headroom under the ceiling. """ if active_bytes + cache_bytes >= ceiling_bytes: return "memory_ceiling" @@ -130,15 +133,21 @@ class _CaptureWatchdog: """Cooperatively stop a capture-run watchdog thread once generation finishes. `policy` states the limits the thread enforced (ceiling, cache bound, polling cadence, - wall budget). + wall budget); `observed` holds the high-water marks it saw (active plus cache, and the + cache term alone) plus the sample count. """ def __init__( - self, stop_event: threading.Event, thread: threading.Thread, policy: dict[str, object] + self, + stop_event: threading.Event, + thread: threading.Thread, + policy: dict[str, object], + observed: dict[str, object], ) -> None: self._stop_event = stop_event self._thread = thread self.policy = policy + self.observed = observed def stop(self) -> None: self._stop_event.set() @@ -198,6 +207,11 @@ def _install_capture_watchdog( f"cache bound {cache_limit_bytes / 1024**3:.1f} GiB, poll {interval_s} s, " f"wall {wall_budget_s:.0f} s" ) + observed: dict[str, object] = { + "peak_total_memory_bytes": 0, + "peak_cache_memory_bytes": 0, + "samples": 0, + } stop_event = threading.Event() started = time.monotonic() abort_path = _abort_artifact_path(variant, out_dir) @@ -222,6 +236,15 @@ def _watch() -> None: stop_event=stop_event, ) continue + observed["samples"] = int(observed["samples"]) + 1 # type: ignore[call-overload] + observed["peak_total_memory_bytes"] = max( + int(observed["peak_total_memory_bytes"]), # type: ignore[call-overload] + active_bytes + cache_bytes, + ) + observed["peak_cache_memory_bytes"] = max( + int(observed["peak_cache_memory_bytes"]), # type: ignore[call-overload] + cache_bytes, + ) reason = _watchdog_breach_reason( active_bytes=active_bytes, cache_bytes=cache_bytes, @@ -248,7 +271,7 @@ def _watch() -> None: thread = threading.Thread(target=_watch, name=f"{variant}-capture-watchdog", daemon=True) thread.start() - return _CaptureWatchdog(stop_event, thread, policy) + return _CaptureWatchdog(stop_event, thread, policy, observed) def _install_memory_caps() -> None: @@ -539,6 +562,7 @@ def main(argv: list[str] | None = None) -> int: ) finally: watchdog.stop() + print(f"watchdog[{args.variant}]: observed {json.dumps(watchdog.observed)}") # Filename uses underscore separator (filesystem-safe) regardless of variant naming. safe_name = args.variant.replace("-", "_") diff --git a/scripts/ab_taef2_bn_domain.py b/scripts/ab_taef2_bn_domain.py index bda2a32..3562526 100644 --- a/scripts/ab_taef2_bn_domain.py +++ b/scripts/ab_taef2_bn_domain.py @@ -40,9 +40,10 @@ CONDITIONS = ("vanilla_vae", "bn_inverse", "identity") _TAEF2_CONDITIONS = ("bn_inverse", "identity") _WORKER_TIMEOUT_S = {"vanilla_vae": 1500, "bn_inverse": 300, "identity": 300} -# Every worker bounds MLX's retained-buffer pool (its default limit sits near device memory); -# the watchdog counts the pool toward its ceiling, so the full-VAE arm's 6 GiB is what keeps -# 28 GiB reachable on 32 GB while still holding that decode's transient buffers. +# Every worker bounds MLX's retained-buffer pool through the watchdog, which counts the pool +# toward its ceiling and records the bound. The full-VAE arm runs under a 12 GB wired cap and +# peaks near 2.8 GiB, so 6 GiB holds that decode's freed transients with room to spare; the +# TAEF2 arms need far less. _CACHE_LIMIT_BYTES = { "vanilla_vae": 6 * 1024**3, "bn_inverse": 2 * 1024**3, @@ -215,6 +216,8 @@ def _worker_main(args: argparse.Namespace) -> int: "unit_wall_s": round(time.perf_counter() - started, 3), "process_peak_memory_gb": round(peak_gb, 3), "installed_cap_gb": installed_cap_gb, + "watchdog": watchdog.policy, + "watchdog_observed": watchdog.observed, }, indent=2, ) diff --git a/scripts/bench_decode.py b/scripts/bench_decode.py index bf13344..8b0bd38 100644 --- a/scripts/bench_decode.py +++ b/scripts/bench_decode.py @@ -128,10 +128,11 @@ def _parse_worker_stdout(stdout: str) -> dict[str, Any]: return payload -# The decode-rep worker's cache bound. It must cover the heaviest rep's peak (the FLUX.1 VAE -# at 512x512 reaches ~3.7 GiB) so nothing is evicted between the warmup and the timed decode, -# or the committed steady-state timings would move; 28 GiB ceiling - 4 GiB pool still leaves -# every rep far below the memory arm. +# The decode-rep worker's cache bound. After the untimed warmup the pool holds that decode's +# freed transients, and the timed decode reuses them; the pool cannot exceed the rep's peak +# active (the heaviest, the FLUX.1 VAE at 512x512, reaches 3.7 GiB including its resident +# weights), so a bound at 4 GiB keeps the warm pool intact and the committed steady-state +# timings comparable. Every rep sits far below the 28 GiB memory arm either way. _BENCH_CACHE_LIMIT_BYTES = 4 * 1024**3 @@ -203,16 +204,13 @@ def _run_one_rep( except (OSError, json.JSONDecodeError): abort_payload = None if isinstance(abort_payload, dict) and abort_payload.get("status") == "aborted": + from scripts.run_showcase import _describe_watchdog_abort + return { "condition": condition, "rep": rep, "status": "failed", - "error": ( - f"watchdog aborted: {abort_payload.get('reason', 'unknown')} " - f"(active={abort_payload.get('active_memory_bytes')} bytes, " - f"ceiling={abort_payload.get('ceiling_bytes')} bytes, " - f"elapsed={abort_payload.get('elapsed_s')}s)" - ), + "error": f"watchdog aborted: {_describe_watchdog_abort(abort_payload)}", } # Cap rejected at startup, OOM, jetsam, etc. Parse stderr for hints. return { @@ -282,10 +280,21 @@ def _run_orchestrator( installed_caps = sorted( {r.get("installed_cap_gb") for r in successes if r.get("installed_cap_gb") is not None} ) + # The watchdog policy is identical across reps of one condition (same worker code path); + # the observed active+cache peak is per rep, like peak_memory_gb. + watchdog_policies = [r["watchdog"] for r in successes if r.get("watchdog") is not None] + per_rep_peak_total = [ + r["watchdog_observed"]["peak_total_memory_bytes"] / 1024**3 + for r in successes + if r.get("watchdog_observed") is not None + ] return { "condition": condition, "applied_cap_gb": cap_gb, "installed_cap_gb": installed_caps[0] if len(installed_caps) == 1 else installed_caps, + "watchdog": watchdog_policies[0] if watchdog_policies else None, + "per_rep_peak_total_memory_gb": per_rep_peak_total, + "max_peak_total_memory_gb": max(per_rep_peak_total) if per_rep_peak_total else None, "reps": len(successes), "per_rep_seconds": per_rep_seconds, "median_seconds": statistics.median(per_rep_seconds), @@ -530,6 +539,7 @@ def _worker_main(args: argparse.Namespace) -> int: "requested_cap_gb": args.applied_cap_gb, "installed_cap_gb": installed_cap_gb, "watchdog": watchdog.policy, + "watchdog_observed": watchdog.observed, } ) ) diff --git a/scripts/run_showcase.py b/scripts/run_showcase.py index a9eab9d..a8c3f55 100644 --- a/scripts/run_showcase.py +++ b/scripts/run_showcase.py @@ -445,12 +445,18 @@ def _validate_live_artifacts( _LIVE_SCENARIOS = frozenset({"live_preview", "zimage_live_preview", "combined"}) _LIVE_WALL_BUDGET_S = 3300.0 _MEMORY_HEADROOM_BYTES = 4 * 1024**3 -# Every model-loading worker bounds MLX's retained-buffer pool. Freed buffers sit in that pool, -# resident but not "active", and its default limit is near device memory; the watchdog counts -# the pool toward the ceiling, so the bound is what makes the arithmetic close on 32 GB: -# the heaviest live scenario (Z-Image) peaks near 25.9 GiB active, and 25.9 + 2 stays under -# the 28 GiB ceiling. The decode-rep worker uses a larger bound (see scripts/bench_decode.py). -_LIVE_CACHE_LIMIT_BYTES = 2 * 1024**3 +# Every model-loading worker bounds MLX's retained-buffer pool and the watchdog counts that +# pool toward the ceiling. Under the harness caps the memory limit (22 GiB, or cap + 2) already +# drains the pool whenever an allocation would cross it (mlx 0.32.2: cached buffers are +# released before a miss that would exceed the limit), so the resident footprint never exceeds +# max(memory limit, peak active); the explicit bound is what the report can state, and it caps +# the pool during the low-active phases (model load, decode) where the memory limit is far +# away. Measured 2026-09-18 (mlx 0.32.2, M1 Max, live_preview at 512x512, three interleaved +# reps each): with the bound at 4 GiB the pool tops out at ~4.1 GiB and the loop takes a +# median 10.76 s; with it at 20 GiB the pool grows to ~17 GiB and the loop takes 11.07 s, so +# the bound is neutral for the committed live wall clocks. `combined` behaved the same +# (8.45 s vs 8.84 s, single runs). The decode-rep worker uses its own bound (bench_decode.py). +_LIVE_CACHE_LIMIT_BYTES = 4 * 1024**3 _WATCHDOG_INTERVAL_S = 0.05 @@ -654,7 +660,9 @@ def _live_watchdog_breach_reason( The memory arm compares the MLX allocator's resident footprint, active plus retained cache, against the ceiling: cached buffers are freed from the graph's point of view but - still occupy unified memory until MLX releases them (`mx.clear_cache` or the cache limit). + still occupy unified memory until MLX releases them (`mx.clear_cache`, the cache limit, or + a memory-limit reclaim). Non-MLX process memory (torch for LPIPS, numpy, the Python heap) + is outside both counters and covered only by the 4 GiB headroom under the ceiling. """ if active_bytes + cache_bytes >= ceiling_bytes: return "memory_ceiling" @@ -663,6 +671,24 @@ def _live_watchdog_breach_reason( return None +def _describe_watchdog_abort(payload: dict[str, Any]) -> str: + """Describe a watchdog abort payload in one operator-facing line. + + Names the reason, every memory term the watchdog compared, and the error text when the + abort came from a failed sample. + """ + parts = [ + f"active={payload.get('active_memory_bytes')} bytes", + f"cache={payload.get('cache_memory_bytes')} bytes", + f"total={payload.get('total_memory_bytes')} bytes", + f"ceiling={payload.get('ceiling_bytes')} bytes", + f"elapsed={payload.get('elapsed_s')}s", + ] + if payload.get("error") is not None: + parts.append(f"error={payload['error']}") + return f"{payload.get('reason', 'unknown')} ({', '.join(parts)})" + + def _commit_watchdog_abort( result_path: Path, payload: dict[str, Any], *, stop_event: threading.Event ) -> None: @@ -696,15 +722,22 @@ class _LiveWatchdog: """Cooperatively stop a live-worker watchdog thread after generation. `policy` states the limits the thread enforced (ceiling, cache bound, polling cadence, - wall budget) so a worker can record them next to its result. + wall budget); `observed` holds the high-water marks the thread saw (active plus cache, + and the cache term alone) plus the sample count, so a worker can record how close it + came next to its result. """ def __init__( - self, stop_event: threading.Event, thread: threading.Thread, policy: dict[str, Any] + self, + stop_event: threading.Event, + thread: threading.Thread, + policy: dict[str, Any], + observed: dict[str, Any], ) -> None: self._stop_event = stop_event self._thread = thread self.policy = policy + self.observed = observed def stop(self) -> None: self._stop_event.set() @@ -739,6 +772,11 @@ def _install_live_watchdog( "interval_s": interval_s, "wall_budget_s": wall_budget_s, } + observed: dict[str, Any] = { + "peak_total_memory_bytes": 0, + "peak_cache_memory_bytes": 0, + "samples": 0, + } stop_event = threading.Event() started = time.monotonic() @@ -762,6 +800,13 @@ def _watch() -> None: stop_event=stop_event, ) continue + observed["samples"] += 1 + observed["peak_total_memory_bytes"] = max( + observed["peak_total_memory_bytes"], active_bytes + cache_bytes + ) + observed["peak_cache_memory_bytes"] = max( + observed["peak_cache_memory_bytes"], cache_bytes + ) reason = _live_watchdog_breach_reason( active_bytes=active_bytes, cache_bytes=cache_bytes, @@ -788,7 +833,7 @@ def _watch() -> None: thread = threading.Thread(target=_watch, name=f"{scenario}-watchdog", daemon=True) thread.start() - return _LiveWatchdog(stop_event, thread, policy) + return _LiveWatchdog(stop_event, thread, policy, observed) # --------------------------------------------------------------------------- @@ -972,10 +1017,7 @@ def _run_live_scenario_subprocess(scenario: str, args: argparse.Namespace) -> di partial = None if isinstance(partial, dict) and partial.get("status") == "aborted": raise TaefError( - f"live scenario worker {scenario!r} aborted: {partial.get('reason', 'unknown')} " - f"(active={partial.get('active_memory_bytes')} bytes, " - f"ceiling={partial.get('ceiling_bytes')} bytes, " - f"elapsed={partial.get('elapsed_s')}s)" + f"live scenario worker {scenario!r} aborted: {_describe_watchdog_abort(partial)}" ) raise TaefError( f"live scenario worker {scenario!r} failed with exit {completed.returncode}: " @@ -1096,6 +1138,7 @@ def main(argv: list[str] | None = None) -> int: finally: watchdog.stop() result["watchdog"] = watchdog.policy + result["watchdog_observed"] = watchdog.observed _write_report(args.live_result, result) return 0 diff --git a/tests/test_ab_taef2_bn_domain.py b/tests/test_ab_taef2_bn_domain.py index 18c7bbf..60d057d 100644 --- a/tests/test_ab_taef2_bn_domain.py +++ b/tests/test_ab_taef2_bn_domain.py @@ -6,6 +6,7 @@ import json from pathlib import Path +from typing import ClassVar import mlx.core as mx import numpy as np @@ -93,6 +94,23 @@ def test_pending_units_reruns_results_that_belong_to_a_different_latent(tmp_path assert ab._pending_units(tmp_path, "lat", latent_sha256="bbbb") == list(ab.CONDITIONS) +def test_install_worker_limits_maps_each_condition_to_its_wired_cap(monkeypatch) -> None: + """Catches: the full-VAE arm (the heaviest decode) running under TAEF2's ~1-2 GB wired cap + because the condition -> bench-condition mapping collapsed to "taef2".""" + import scripts.ab_taef2_bn_domain as ab + import scripts.bench_decode as bench + + caps: list[int] = [] + monkeypatch.setattr(bench, "_install_memory_caps", lambda cap: caps.append(cap) or cap) + + assert ab._install_worker_limits("vanilla_vae") == bench._resolve_cap_gb( + condition="vanilla_vae" + ) + for condition in ab._TAEF2_CONDITIONS: + assert ab._install_worker_limits(condition) == bench._resolve_cap_gb(condition="taef2") + assert caps[0] > caps[1] == caps[2] + + def test_every_worker_hands_its_cache_bound_to_the_watchdog(monkeypatch, tmp_path) -> None: """Catches: a worker (the full-VAE arm is the heaviest) running with MLX's near-device-size default cache limit, or with the watchdog's generic bound instead of the condition's own @@ -131,6 +149,43 @@ def _fake_install(result_path, scenario: str, **kwargs: object) -> _FakeWatchdog assert 0 < ab._cache_limit_bytes(condition) <= 8 * 1024**3 +def test_unit_result_records_the_watchdog_policy_and_observed_peak(monkeypatch, tmp_path) -> None: + """Catches: an A/B unit result that states the wired cap it ran under but not the cache + bound or the observed active+cache peak, so the report cannot say how close it came.""" + import argparse + + import scripts.ab_taef2_bn_domain as ab + import scripts.run_showcase as rs + + class _FakeWatchdog: + policy: ClassVar[dict[str, object]] = {"cache_limit_bytes": 2 * 1024**3, "interval_s": 0.05} + observed: ClassVar[dict[str, object]] = { + "peak_total_memory_bytes": 5, + "peak_cache_memory_bytes": 3, + } + + def stop(self) -> None: + pass + + monkeypatch.setattr(rs, "_install_live_watchdog", lambda *a, **kw: _FakeWatchdog()) + monkeypatch.setattr(ab, "_install_worker_limits", lambda condition: 7) + monkeypatch.setattr( + "scripts.bench_decode._prep_full_vae_flux2", + lambda latent, h, w: lambda: mx.zeros((1, 2, 2, 3), dtype=mx.uint8), + ) + monkeypatch.setattr(ab, "_save_png", lambda image, target: target.write_bytes(b"png")) + latent = tmp_path / "lat.safetensors" + mx.save_safetensors( + str(latent), {"latent": mx.zeros((1, 4, 16)), "height": mx.array(16), "width": mx.array(16)} + ) + args = argparse.Namespace(worker="vanilla_vae", latent=latent, out_dir=tmp_path) + + assert ab._worker_main(args) == 0 + result = json.loads(ab._unit_result_path(tmp_path, "lat", "vanilla_vae").read_text()) + assert result["watchdog"] == _FakeWatchdog.policy + assert result["watchdog_observed"] == _FakeWatchdog.observed + + def test_score_keeps_a_computed_lpips_when_the_other_arm_fails(tmp_path: Path, monkeypatch) -> None: """Catches: one arm's scoring error silently discarding the other arm's valid LPIPS.""" import scripts.ab_taef2_bn_domain as ab diff --git a/tests/test_bench_decode.py b/tests/test_bench_decode.py index 4d88a8e..0d8d782 100644 --- a/tests/test_bench_decode.py +++ b/tests/test_bench_decode.py @@ -350,6 +350,7 @@ def test_worker_main_clears_stale_watchdog_abort_artifact_before_running( class _FakeWatchdog: policy: ClassVar[dict[str, object]] = {"cache_limit_bytes": 0, "interval_s": 0.05} + observed: ClassVar[dict[str, object]] = {} def stop(self) -> None: pass @@ -401,6 +402,7 @@ def _fake_decode() -> str: class _FakeWatchdog: policy: ClassVar[dict[str, object]] = {"cache_limit_bytes": 0, "interval_s": 0.05} + observed: ClassVar[dict[str, object]] = {} def stop(self) -> None: pass @@ -469,6 +471,7 @@ def test_worker_main_installs_watchdog_with_condition_scoped_wall_budget( class _FakeWatchdog: policy: ClassVar[dict[str, object]] = {"cache_limit_bytes": 0, "interval_s": 0.05} + observed: ClassVar[dict[str, object]] = {} def stop(self) -> None: pass @@ -547,6 +550,7 @@ def test_worker_main_passes_a_measurement_preserving_cache_bound_to_the_watchdog class _FakeWatchdog: policy: ClassVar[dict[str, object]] = {"cache_limit_bytes": 4 * 1024**3, "interval_s": 0.05} + observed: ClassVar[dict[str, object]] = {} def stop(self) -> None: pass @@ -574,7 +578,10 @@ def _fake_install(result_path: Path, scenario: str, **kwargs: object) -> _FakeWa assert bench._worker_main(args) == 0 assert install_calls[0]["cache_limit_bytes"] == bench._BENCH_CACHE_LIMIT_BYTES - assert 4 * 1024**3 <= bench._BENCH_CACHE_LIMIT_BYTES <= 8 * 1024**3 + # The heaviest rep in the committed report peaks at 3.7 GiB active; the pool of freed + # transients after its warmup cannot exceed that, so a bound at or above it keeps the pool. + assert bench._BENCH_CACHE_LIMIT_BYTES >= 3.7 * 1024**3 + assert bench._BENCH_CACHE_LIMIT_BYTES <= 8 * 1024**3 def test_worker_main_records_the_watchdog_policy_in_the_sentinel( @@ -595,6 +602,7 @@ def test_worker_main_records_the_watchdog_policy_in_the_sentinel( class _FakeWatchdog: policy: ClassVar[dict[str, object]] = {"cache_limit_bytes": 4 * 1024**3, "interval_s": 0.05} + observed: ClassVar[dict[str, object]] = {"peak_total_memory_bytes": 7} def stop(self) -> None: pass @@ -618,3 +626,88 @@ def stop(self) -> None: assert bench._worker_main(args) == 0 payload = bench._parse_worker_stdout(capsys.readouterr().out) assert payload["watchdog"] == {"cache_limit_bytes": 4 * 1024**3, "interval_s": 0.05} + assert payload["watchdog_observed"] == {"peak_total_memory_bytes": 7} + + +def test_run_orchestrator_carries_the_watchdog_policy_and_observed_peaks_into_the_report( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """Catches: `_run_orchestrator` rebuilding the condition block from named keys and dropping the + per-rep watchdog policy and observed active+cache peak, so the committed report never + learns how close the reps came to the ceiling.""" + import scripts.bench_decode as bench + + reps = [ + { + "status": "ok", + "elapsed_s": 0.1 * (i + 1), + "peak_memory_gb": 1.0, + "installed_cap_gb": 4, + "image_path": "x.webp", + "watchdog": {"cache_limit_bytes": 4 * 1024**3, "interval_s": 0.05}, + "watchdog_observed": {"peak_total_memory_bytes": (i + 1) * 1024**3}, + } + for i in range(3) + ] + monkeypatch.setattr(bench, "_run_one_rep", lambda **kw: reps[kw["rep"]]) + + block = bench._run_orchestrator( + latent_path=tmp_path / "l.safetensors", condition="taef1", reps=3, save_dir=tmp_path + ) + + assert block["watchdog"] == {"cache_limit_bytes": 4 * 1024**3, "interval_s": 0.05} + assert block["per_rep_peak_total_memory_gb"] == [1.0, 2.0, 3.0] + assert block["max_peak_total_memory_gb"] == 3.0 + + +def test_run_one_rep_abort_message_names_cache_total_and_error(tmp_path: Path, monkeypatch) -> None: + """Catches: `_run_one_rep` formatting only active and ceiling from the abort artifact, so + a cache-driven breach reads as a non-breach and a `sample_error` never shows its cause.""" + import subprocess + + from scripts.bench_decode import _run_one_rep, _watchdog_abort_path + + save_to = tmp_path / "rep0.webp" + abort_path = _watchdog_abort_path(save_to) + monkeypatch.setattr( + subprocess, + "run", + lambda *a, **kw: subprocess.CompletedProcess(args=[], returncode=70, stdout="", stderr=""), + ) + abort_path.write_text( + json.dumps( + { + "status": "aborted", + "reason": "memory_ceiling", + "active_memory_bytes": 20, + "cache_memory_bytes": 9, + "total_memory_bytes": 29, + "ceiling_bytes": 28, + "elapsed_s": 1.5, + } + ) + ) + result = _run_one_rep( + latent_path=tmp_path / "l.safetensors", + condition="taef1", + flux_variant="flux1-dev", + rep=0, + save_to=save_to, + cap_gb=1, + ) + assert result["status"] == "failed" + assert "cache=9 bytes" in result["error"] + assert "total=29 bytes" in result["error"] + + abort_path.write_text( + json.dumps({"status": "aborted", "reason": "sample_error", "error": "RuntimeError: gone"}) + ) + result = _run_one_rep( + latent_path=tmp_path / "l.safetensors", + condition="taef1", + flux_variant="flux1-dev", + rep=0, + save_to=save_to, + cap_gb=1, + ) + assert "RuntimeError: gone" in result["error"] diff --git a/tests/test_capture_latent.py b/tests/test_capture_latent.py index 7749b90..203cb53 100644 --- a/tests/test_capture_latent.py +++ b/tests/test_capture_latent.py @@ -354,10 +354,39 @@ def test_capture_watchdog_bounds_the_cache_pool_and_records_its_policy( watchdog.stop() assert limits == [cl._CAPTURE_CACHE_LIMIT_BYTES] - assert 0 < cl._CAPTURE_CACHE_LIMIT_BYTES <= 2 * 1024**3 + # Captures run the live generation recipes, so they share the live bound's floor. + from scripts.run_showcase import _LIVE_CACHE_LIMIT_BYTES + + assert cl._CAPTURE_CACHE_LIMIT_BYTES == _LIVE_CACHE_LIMIT_BYTES assert watchdog.policy == { "ceiling_bytes": 28 * 1024**3, "cache_limit_bytes": cl._CAPTURE_CACHE_LIMIT_BYTES, "interval_s": 0.05, "wall_budget_s": cl._WALL_BUDGET_S, } + + +def test_capture_watchdog_reports_the_observed_active_plus_cache_peak( + monkeypatch, tmp_path: Path +) -> None: + """Catches: a capture run that prints its active peak only while the watchdog's own + reading, active plus cache, came within a few MiB of firing; the high-water marks of the + sum and of the cache term must both track the samples.""" + import time + + import scripts._capture_latent as cl + + active = iter([1, 5, 2]) + cache = iter([1, 3, 6]) + monkeypatch.setattr(cl.mx, "device_info", lambda: {"memory_size": 32 * 1024**3}) + monkeypatch.setattr(cl.mx, "set_cache_limit", lambda n: 0) + monkeypatch.setattr(cl.mx, "get_active_memory", lambda: next(active, 2)) + monkeypatch.setattr(cl.mx, "get_cache_memory", lambda: next(cache, 1)) + + watchdog = cl._install_capture_watchdog("flux1-dev", tmp_path, interval_s=0.002) + time.sleep(0.1) + watchdog.stop() + + assert watchdog.observed["peak_total_memory_bytes"] == 8 + assert watchdog.observed["peak_cache_memory_bytes"] == 6 + assert watchdog.observed["samples"] >= 3 diff --git a/tests/test_run_showcase.py b/tests/test_run_showcase.py index 2305706..5007742 100644 --- a/tests/test_run_showcase.py +++ b/tests/test_run_showcase.py @@ -1,6 +1,7 @@ """Plumbing tests for scripts/run_showcase.py.""" import json +import time from pathlib import Path from typing import ClassVar @@ -106,6 +107,7 @@ def test_live_worker_mode_runs_one_raw_scenario(tmp_path: Path, monkeypatch) -> class _FakeWatchdog: policy: ClassVar[dict[str, object]] = {"cache_limit_bytes": 1, "interval_s": 0.05} + observed: ClassVar[dict[str, object]] = {"peak_total_memory_bytes": 9} def stop(self) -> None: watchdog_events.append("stopped") @@ -142,6 +144,7 @@ def _fake_install(result_path: Path, scenario: str) -> _FakeWatchdog: assert json.loads(result_path.read_text()) == { "status": "ok", "watchdog": {"cache_limit_bytes": 1, "interval_s": 0.05}, + "watchdog_observed": {"peak_total_memory_bytes": 9}, } @@ -304,36 +307,71 @@ def test_live_watchdog_aborts_explicitly_when_a_memory_sample_fails( def test_live_watchdog_bounds_the_cache_pool_before_polling_and_records_its_policy( monkeypatch, tmp_path: Path ) -> None: - """Catches: a live worker running with MLX's near-device-size default cache limit (the - accounting can then never close on 32 GB), and a policy the report cannot reproduce. - The bound is installed by the watchdog so every model-loading harness path gets it.""" + """Catches: a live worker running with MLX's default cache limit (the policy the report + states would then be a fiction), the bound installed after the first sample, and a + policy the report cannot reproduce. The bound is installed by the watchdog so every + model-loading harness path gets it from one place.""" import mlx.core as mx import scripts.run_showcase as rs - limits: list[int] = [] + events: list[str] = [] monkeypatch.setattr(mx, "device_info", lambda: {"memory_size": 32 * 1024**3}) - monkeypatch.setattr(mx, "set_cache_limit", lambda n: limits.append(n) or 0) - monkeypatch.setattr(mx, "get_active_memory", lambda: 0) + monkeypatch.setattr(mx, "set_cache_limit", lambda n: events.append(f"cache:{n}") or 0) + monkeypatch.setattr(mx, "get_active_memory", lambda: events.append("sample") or 0) monkeypatch.setattr(mx, "get_cache_memory", lambda: 0) - watchdog = rs._install_live_watchdog(tmp_path / "r.json", "scn", cache_limit_bytes=3 * 1024**3) + watchdog = rs._install_live_watchdog( + tmp_path / "r.json", "scn", cache_limit_bytes=3 * 1024**3, interval_s=0.005 + ) + time.sleep(0.05) watchdog.stop() - assert limits == [3 * 1024**3] + assert events[0] == f"cache:{3 * 1024**3}" + assert "sample" in events[1:] assert watchdog.policy == { "ceiling_bytes": 28 * 1024**3, "cache_limit_bytes": 3 * 1024**3, - "interval_s": 0.05, + "interval_s": 0.005, "wall_budget_s": rs._LIVE_WALL_BUDGET_S, } -def test_live_watchdog_default_cache_bound_closes_the_ceiling_arithmetic() -> None: - """Catches: a default cache bound so large that a live worker's measured peak plus a - full cache pool overshoots the 28 GiB ceiling on 32 GB (Z-Image peaks at ~25.9 GiB).""" +def test_live_watchdog_reports_the_observed_active_plus_cache_peak( + monkeypatch, tmp_path: Path +) -> None: + """Catches: a result whose `peak_memory_gb` (active only) says 2 GiB of headroom while the + watchdog's own reading, active plus cache, came within a few MiB of firing. The watchdog + keeps a high-water mark of the sum it compares, and of the cache term alone.""" + import mlx.core as mx + import scripts.run_showcase as rs + + active = iter([1, 5, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2]) + cache = iter([1, 3, 6, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1]) + monkeypatch.setattr(mx, "device_info", lambda: {"memory_size": 32 * 1024**3}) + monkeypatch.setattr(mx, "set_cache_limit", lambda n: 0) + monkeypatch.setattr(mx, "get_active_memory", lambda: next(active, 2)) + monkeypatch.setattr(mx, "get_cache_memory", lambda: next(cache, 1)) + + watchdog = rs._install_live_watchdog(tmp_path / "r.json", "scn", interval_s=0.002) + time.sleep(0.1) + watchdog.stop() + + assert watchdog.observed == { + "peak_total_memory_bytes": 8, + "peak_cache_memory_bytes": 6, + "samples": watchdog.observed["samples"], + } + assert watchdog.observed["samples"] >= 3 + + +def test_live_cache_bound_is_not_below_the_value_the_timing_check_found_neutral() -> None: + """Catches: the live bound shrinking below 4 GiB, the only value the 2026-09-18 timing + check compared against an unbounded pool (live_preview 10.76 s vs 11.07 s, three reps + each); a smaller pool is unmeasured and could move the committed live wall clocks. This + pins the constant to the measured floor; it cannot verify the measurement itself.""" import scripts.run_showcase as rs - assert 0 < rs._LIVE_CACHE_LIMIT_BYTES <= 2 * 1024**3 + assert rs._LIVE_CACHE_LIMIT_BYTES >= 4 * 1024**3 def test_json_schema_rejects_unknown_version(tmp_path: Path) -> None: @@ -758,6 +796,7 @@ def test_vs_vae_worker_installs_active_memory_watchdog( class _FakeWatchdog: policy: ClassVar[dict[str, object]] = {"cache_limit_bytes": 0, "interval_s": 0.05} + observed: ClassVar[dict[str, object]] = {} def stop(self) -> None: watchdog_events.append("stopped") @@ -1243,3 +1282,43 @@ def _git(*argv: str) -> None: (tmp_path / "src" / "a.py").write_text("x = 2\n") assert rs._detect_source_version().endswith("-dirty") + + +def test_live_scenario_abort_message_names_cache_total_and_error( + tmp_path: Path, monkeypatch +) -> None: + """Catches: the orchestrator's abort message formatting only active and ceiling, so a + cache-driven breach reads as a non-breach and a `sample_error` never shows its cause.""" + import argparse + import subprocess + + import scripts.run_showcase as rs + + from mlx_taef.errors import TaefError + + report = tmp_path / "report.json" + args = argparse.Namespace(report=report, cap_gb=None) + result_path = tmp_path / ".report-partials" / "live_preview.json" + + def _fake_run(*a: object, **kwargs: object) -> subprocess.CompletedProcess[str]: + result_path.write_text(json.dumps(payload)) + return subprocess.CompletedProcess(args=[], returncode=70, stdout="", stderr="") + + monkeypatch.setattr(rs.subprocess, "run", _fake_run) + + payload = { + "status": "aborted", + "reason": "memory_ceiling", + "active_memory_bytes": 20, + "cache_memory_bytes": 9, + "total_memory_bytes": 29, + "ceiling_bytes": 28, + "elapsed_s": 1.5, + } + with pytest.raises(TaefError, match="cache=9 bytes") as excinfo: + rs._run_live_scenario_subprocess("live_preview", args) + assert "total=29 bytes" in str(excinfo.value) + + payload = {"status": "aborted", "reason": "sample_error", "error": "RuntimeError: gone"} + with pytest.raises(TaefError, match="RuntimeError: gone"): + rs._run_live_scenario_subprocess("live_preview", args) From 3018f922643ecf2c2f4a055b6386ba845f3e6d9a Mon Sep 17 00:00:00 2001 From: Denis Ineshin Date: Fri, 18 Sep 2026 22:39:52 +0200 Subject: [PATCH 3/3] Print the wall budget in abort lines and wait for samples in the thread tests The operator line for a watchdog abort now lists every term the artifact carries, including the wall budget, and skips terms that are absent instead of printing None. The thread tests wait for the sample count they need rather than sleeping a fixed interval. The CHANGELOG no longer calls the two allocator counters a resident figure. --- CHANGELOG.md | 6 +++--- scripts/run_showcase.py | 24 +++++++++++++++--------- tests/test_capture_latent.py | 5 ++++- tests/test_run_showcase.py | 27 ++++++++++++++++++++++++--- 4 files changed, 46 insertions(+), 16 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 69ec95f..3fc5351 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,8 +9,8 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Changed - The showcase, benchmark, latent-capture and A/B harness workers now count MLX's retained - buffer cache toward their memory ceiling instead of active memory alone, which is the number - MLX itself reports as resident. Each watchdog bounds the cache pool when it starts and states + buffer cache toward their memory ceiling instead of active memory alone; together those are + the bytes the MLX allocator holds. Each watchdog bounds the cache pool when it starts and states that bound, its polling cadence (now 50 ms), the ceiling and the wall budget in its abort artifact and in each worker's result, together with the highest active-plus-cache reading it saw, so a report row can say how close a run came to the ceiling rather than only its active @@ -18,7 +18,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 so the timed decode still runs from a warm pool) and the A/B script's existing per-condition values. A memory sample that raises now aborts the worker with reason `sample_error` and the error text rather than ending the watchdog thread in silence, and the orchestrators print - every term the watchdog compared. Harness only: the library and its decode path are + every term the abort artifact carries. Harness only: the library and its decode path are unchanged, and the committed showcase numbers were not re-measured. ## [0.8.2] - 2026-09-17 diff --git a/scripts/run_showcase.py b/scripts/run_showcase.py index a8c3f55..1a562a1 100644 --- a/scripts/run_showcase.py +++ b/scripts/run_showcase.py @@ -674,18 +674,24 @@ def _live_watchdog_breach_reason( def _describe_watchdog_abort(payload: dict[str, Any]) -> str: """Describe a watchdog abort payload in one operator-facing line. - Names the reason, every memory term the watchdog compared, and the error text when the - abort came from a failed sample. + Names the reason, every term the watchdog compared that the payload carries (memory + terms, ceiling, elapsed time, wall budget), and the error text when the abort came from + a failed sample. Terms absent from the payload are left out rather than printed as None. """ + fields = ( + ("active", "active_memory_bytes", " bytes"), + ("cache", "cache_memory_bytes", " bytes"), + ("total", "total_memory_bytes", " bytes"), + ("ceiling", "ceiling_bytes", " bytes"), + ("elapsed", "elapsed_s", "s"), + ("wall_budget", "wall_budget_s", "s"), + ("error", "error", ""), + ) parts = [ - f"active={payload.get('active_memory_bytes')} bytes", - f"cache={payload.get('cache_memory_bytes')} bytes", - f"total={payload.get('total_memory_bytes')} bytes", - f"ceiling={payload.get('ceiling_bytes')} bytes", - f"elapsed={payload.get('elapsed_s')}s", + f"{label}={payload[key]}{unit}" + for label, key, unit in fields + if payload.get(key) is not None ] - if payload.get("error") is not None: - parts.append(f"error={payload['error']}") return f"{payload.get('reason', 'unknown')} ({', '.join(parts)})" diff --git a/tests/test_capture_latent.py b/tests/test_capture_latent.py index 203cb53..22cb6ff 100644 --- a/tests/test_capture_latent.py +++ b/tests/test_capture_latent.py @@ -384,7 +384,10 @@ def test_capture_watchdog_reports_the_observed_active_plus_cache_peak( monkeypatch.setattr(cl.mx, "get_cache_memory", lambda: next(cache, 1)) watchdog = cl._install_capture_watchdog("flux1-dev", tmp_path, interval_s=0.002) - time.sleep(0.1) + end = time.monotonic() + 5.0 + while int(watchdog.observed["samples"]) < 3: # type: ignore[call-overload] + assert time.monotonic() < end, "watchdog took fewer than 3 samples in 5s" + time.sleep(0.001) watchdog.stop() assert watchdog.observed["peak_total_memory_bytes"] == 8 diff --git a/tests/test_run_showcase.py b/tests/test_run_showcase.py index 5007742..99563f1 100644 --- a/tests/test_run_showcase.py +++ b/tests/test_run_showcase.py @@ -217,6 +217,15 @@ def test_live_watchdog_breach_reason_counts_retained_cache_toward_the_ceiling() ) +def _wait_for_samples(watchdog: object, n: int, *, deadline_s: float = 5.0) -> None: + """Block until the watchdog thread has taken `n` samples (deterministic, no fixed sleep).""" + end = time.monotonic() + deadline_s + while watchdog.observed["samples"] < n: # type: ignore[attr-defined] + if time.monotonic() > end: + raise AssertionError(f"watchdog took fewer than {n} samples in {deadline_s}s") + time.sleep(0.001) + + def _install_live_watchdog_with_fake_mlx( monkeypatch, tmp_path: Path, @@ -323,7 +332,7 @@ def test_live_watchdog_bounds_the_cache_pool_before_polling_and_records_its_poli watchdog = rs._install_live_watchdog( tmp_path / "r.json", "scn", cache_limit_bytes=3 * 1024**3, interval_s=0.005 ) - time.sleep(0.05) + _wait_for_samples(watchdog, 1) watchdog.stop() assert events[0] == f"cache:{3 * 1024**3}" @@ -353,7 +362,7 @@ def test_live_watchdog_reports_the_observed_active_plus_cache_peak( monkeypatch.setattr(mx, "get_cache_memory", lambda: next(cache, 1)) watchdog = rs._install_live_watchdog(tmp_path / "r.json", "scn", interval_s=0.002) - time.sleep(0.1) + _wait_for_samples(watchdog, 3) watchdog.stop() assert watchdog.observed == { @@ -1319,6 +1328,18 @@ def _fake_run(*a: object, **kwargs: object) -> subprocess.CompletedProcess[str]: rs._run_live_scenario_subprocess("live_preview", args) assert "total=29 bytes" in str(excinfo.value) + payload = { + "status": "aborted", + "reason": "wall_budget", + "elapsed_s": 61.0, + "wall_budget_s": 60.0, + } + with pytest.raises(TaefError, match="wall_budget") as excinfo: + rs._run_live_scenario_subprocess("live_preview", args) + assert "wall_budget=60.0s" in str(excinfo.value) + assert "None" not in str(excinfo.value) + payload = {"status": "aborted", "reason": "sample_error", "error": "RuntimeError: gone"} - with pytest.raises(TaefError, match="RuntimeError: gone"): + with pytest.raises(TaefError, match="RuntimeError: gone") as excinfo: rs._run_live_scenario_subprocess("live_preview", args) + assert "None" not in str(excinfo.value)