Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 16 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,22 @@ 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; 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
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 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

TAEF2 previews now track the full FLUX.2 VAE closely.
Expand Down
2 changes: 1 addition & 1 deletion COMPARISON.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
103 changes: 90 additions & 13 deletions scripts/_capture_latent.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,12 @@
# 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 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


def _build_argparser() -> argparse.ArgumentParser:
Expand Down Expand Up @@ -103,22 +109,45 @@ 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 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"
if elapsed_s > wall_budget_s:
return "wall_budget"
return None


class _CaptureWatchdog:
"""Cooperatively stop a capture-run watchdog thread once generation finishes."""
"""Cooperatively stop a capture-run watchdog thread once generation finishes.

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); `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],
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()
Expand Down Expand Up @@ -149,30 +178,76 @@ 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 (`<out-dir>/<variant>.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"
)
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)

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
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,
ceiling_bytes=ceiling_bytes,
elapsed_s=elapsed_s,
wall_budget_s=wall_budget_s,
Expand All @@ -186,16 +261,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, observed)


def _install_memory_caps() -> None:
Expand Down Expand Up @@ -486,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("-", "_")
Expand Down
23 changes: 15 additions & 8 deletions scripts/ab_taef2_bn_domain.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,8 +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, and
# the watchdog samples active memory only.
# 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,
Expand All @@ -65,17 +67,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:
Expand Down Expand Up @@ -170,7 +172,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:
Expand Down Expand Up @@ -211,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,
)
Expand Down
33 changes: 26 additions & 7 deletions scripts/bench_decode.py
Original file line number Diff line number Diff line change
Expand Up @@ -128,6 +128,14 @@ def _parse_worker_stdout(stdout: str) -> dict[str, Any]:
return payload


# 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


def _watchdog_abort_path(save_to: Path) -> Path:
"""Where this rep's watchdog abort artifact would be written, if any.

Expand Down Expand Up @@ -196,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 {
Expand Down Expand Up @@ -275,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),
Expand Down Expand Up @@ -457,7 +473,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.
Expand All @@ -476,6 +492,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
Expand Down Expand Up @@ -521,6 +538,8 @@ 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,
"watchdog_observed": watchdog.observed,
}
)
)
Expand Down
5 changes: 3 additions & 2 deletions scripts/capture_examples.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 `<out-dir>/<variant>.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
Expand Down Expand Up @@ -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.

Expand Down
Loading
Loading