* feat(downloads): Xet fast path + accurate progress (FDL W0–W2)
Make model downloads fast and show accurate downloaded/remaining/speed.
Research confirmed hf-xet already implements the IDM/uGet technique
(content-defined chunking, parallel byte-range gets, dedup, resume), and
the spike found all 25 catalog repos are Xet-backed — so the win is
driving Xet well + accurate progress, not a custom downloader.
W1 — maximize + guarantee Xet:
- pin huggingface_hub>=1.7 + hf-xet>=1.1 (was transitive); no hf_transfer
- drive snapshot_download with explicit tqdm_class + max_workers + endpoint
- opt-in HF_XET_HIGH_PERFORMANCE / HDD sequential-write knobs (default off)
- /system/info reports fast_download {xet_enabled, xet_version, high_perf}
W2 — accurate progress:
- dry_run preflight -> install_plan event (exact total/cached/remaining)
- utils/download_aggregator.py: one overall bar; byte bars (by id) vs the
"Fetching N files" count bar; windowed rate; emits one 'aggregate' event
- frontend overall bar (speed/remaining/ETA), cached-skip, ⚡ fast badge
Known limit (verified live): under Xet+hf_hub 1.7.2 per-file byte bars
never advance/close via tqdm, so mid-download the bar is file-granular and
bytes flush to the exact total on completion. Classic-LFS/mirror repos get
true byte progress (W4).
Drive-by: download.py used os.walk without importing os (latent NameError
in _validate_snapshot_has_weights on every install) — fixed.
Tests: tests/backend/setup/test_download_preflight.py (10). Spike + plan
under .planning/quick/260613-fdl-fast-model-downloads/.
Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
* feat(downloads): opt-in mirror + cancel + docs (FDL W4)
- mirror (FDL-10): snapshot_download(endpoint=) honours prefs hf_endpoint /
env HF_ENDPOINT on preflight + download (per-call, no process-wide env).
Documented as the classic-LFS path (no Xet) for restricted networks.
- cancel (FDL-11): POST /models/install/cancel {repo_id} stops further
retries at the next boundary, emits install_cancelled, clears the cooldown
(cancel is intent, not failure). Frontend treats it as a terminator.
- docs (FDL-12): docs/downloading-models.md (Xet fast path, progress
semantics + byte-speed limitation, opt-in tuning, mirror, cancel,
troubleshooting) + README pointer. Docs-sync rule satisfied.
Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
* docs(planning): model-management v2 cleanup plan (mm2)
GSD plan for cleaning the model-management subsystem: registry unload-on-
switch + per-engine unload() (fixes VRAM leak), model_lifecycle facade,
unified idle/timeout config, bounded cooldowns, sidecar VRAM self-report,
cache-fallback logging. Planning artifact only — no code.
Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
* fix(downloads): reconcile with main's HF_HUB_DISABLE_XET; honest status
Rebasing onto main surfaced that main forces HF_HUB_DISABLE_XET=1 (classic
LFS) because Xet progress bypasses the tqdm hook — the same limitation found
here. Reconcile instead of fight:
- /system/info fast_download now reports runtime truth: xet_installed +
xet_active (installed AND not HF_HUB_DISABLE_XET) + xet_enabled alias. The
⚡ badge only shows when Xet actually runs; startup log says
"downloads: Xet disabled → legacy LFS".
- complete(): clear the rate window before the final flush so crediting the
full size in one step can't emit an absurd instantaneous rate.
- docs/downloading-models.md rewritten: default is legacy LFS for accurate
progress; Xet is opt-in via HF_HUB_DISABLE_XET=0. hf-xet pin stays (ready
for a future Xet progress hook).
W2 (preflight total/remaining + aggregate bar + exact completion) is the
value on either path; W1's "maximize Xet" is dormant by main's design.
Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
* feat(downloads): opt-in segmented multi-connection accelerator (FDL W3)
Since main forces Xet off (HF_HUB_DISABLE_XET=1), the default path is
single-stream legacy LFS — so a segmented downloader is the way to get BOTH
parallel speed and live byte progress.
- services/segmented_download.py: async multi-connection Range downloader for
one file — parallel byte-ranges, resume (.part + manifest), per-segment
short-read truncation guard, optional sha256/etag verify, cancel, and a
single-stream fallback when the server won't range. Auth-safe: the HF
Authorization header is sent only to huggingface.co/hf.co and never
forwarded to a CDN host on redirect (unit-tested).
- dispatch (download.py): opt-in via prefs segmented_downloader / env
OMNIVOICE_SEGMENTED_DOWNLOAD (default off). When on and Xet inactive,
fetches each file into the HF cache mirroring hf_hub_download (blobs +
snapshot symlinks + refs/main), feeding real bytes to the aggregator. Any
failure falls back to snapshot_download — never breaks a correct install.
- fix: complete() was adding a full total on top of accumulated segmented
bytes (2x); now replaces byte bars so the sum is exactly total.
Verified live (accelerator on): real byte progress to ~16.6 MB/s, final
bytes==total, /models installed=True, delete frees correctly.
Tests: test_segmented_download.py (7) + aggregator double-count regression.
Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
* test(downloads): relocate FDL tests to top-level; loop-isolate segmented test
CI runs the full suite, which exposed a pre-existing test-isolation leak:
several tests/backend/** fixtures purge core.*/services.* from sys.modules
under a temp OMNIVOICE_DATA_DIR and never restore, leaving core.config/core.db
bound to a dead temp dir. It only bites when collection order puts a purging
test ahead of a real-DB reader (test_longform_jobs). Adding tests under
tests/backend/setup/ reordered collection and tripped it.
Fix without touching the shared (fragile) fixtures or risking class-identity
breakage from a blanket sys.modules restore:
- move the two FDL test files to top-level tests/ (tests/test_fdl_*.py) so
tests/backend/** collection order is identical to main — longform passes.
- rewrite the segmented test to run each case under asyncio.run() (fresh loop)
instead of asyncio.get_event_loop(), which an earlier async test can leave
closed in the full suite.
Full suite green locally: 1364 passed, 0 failed.
Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
---------
Co-authored-by: mergetest <test@local>
Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
301 lines
12 KiB
Python
301 lines
12 KiB
Python
"""HuggingFace download progress — one monkey-patch, every `hf_hub_download`
|
|
reports bytes downloaded through a central callback.
|
|
|
|
`huggingface_hub` uses tqdm for progress bars; we subclass it, intercept
|
|
`update()` calls, and forward (filename, downloaded_bytes, total_bytes) to
|
|
whatever callback is registered. No changes to calling sites across
|
|
transformers / mlx_whisper / diffusers / accelerate — they all route through
|
|
`hf_hub_download`, which uses the patched tqdm.
|
|
|
|
Usage:
|
|
from utils.hf_progress import install, register_listener, unregister_listener
|
|
|
|
install() # once at app startup
|
|
listener_id = register_listener(lambda ev: print(ev))
|
|
# …models download, listener fires…
|
|
unregister_listener(listener_id)
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import contextvars
|
|
import itertools
|
|
import logging
|
|
import threading
|
|
from typing import Callable, Optional
|
|
|
|
logger = logging.getLogger("omnivoice.hf_progress")
|
|
|
|
# Context-scoped active repo_id. Set in the install/delete handler so every
|
|
# tqdm event fired while a snapshot_download runs can be stamped with the
|
|
# originating repo, letting the frontend route per-file events to the right
|
|
# row instead of heuristically matching filename substrings.
|
|
current_repo_id: contextvars.ContextVar[Optional[str]] = contextvars.ContextVar(
|
|
"omnivoice_hf_progress_repo_id", default=None,
|
|
)
|
|
|
|
# Event shape forwarded to listeners. Typed loosely on purpose — SSE encodes
|
|
# it as JSON so consumers read the dict directly.
|
|
# {
|
|
# "filename": str, # desc on the tqdm bar, usually the HF file path
|
|
# "downloaded": int, # bytes pulled so far
|
|
# "total": int | None, # total bytes or None if unknown
|
|
# "pct": float, # 0.0-1.0 (or 0.0 if total unknown)
|
|
# "phase": "start"|"progress"|"done",
|
|
# }
|
|
ProgressEvent = dict
|
|
Listener = Callable[[ProgressEvent], None]
|
|
|
|
_listeners: dict[int, Listener] = {}
|
|
_listener_lock = threading.Lock()
|
|
_listener_counter = itertools.count(1)
|
|
_installed = False
|
|
_install_lock = threading.Lock()
|
|
|
|
# Set by install() to the TrackedTqdm subclass so call sites can drive it
|
|
# explicitly via snapshot_download(tqdm_class=...) instead of relying solely on
|
|
# the global monkey-patch. Xet feeds bytes into whatever tqdm_class is passed,
|
|
# so this is also the xet-aware progress hook (FDL-02).
|
|
_tracked_tqdm_class: Optional[type] = None
|
|
|
|
|
|
def tracked_tqdm_class() -> Optional[type]:
|
|
"""Return the progress-emitting tqdm subclass (or None if install() hasn't
|
|
run / huggingface_hub's tqdm couldn't be patched). Pass it as
|
|
``snapshot_download(tqdm_class=...)`` to drive progress deterministically."""
|
|
return _tracked_tqdm_class
|
|
|
|
|
|
# Optional sink fed every per-file (repo_id, filename, downloaded, total) byte
|
|
# update — used by utils.download_aggregator to build the overall aggregate bar
|
|
# (FDL-06). Kept as a setter to avoid a circular import (this module must not
|
|
# import the aggregator). Signature: fn(repo_id, filename, downloaded, total).
|
|
_byte_sink: Optional[Callable] = None
|
|
|
|
|
|
def set_byte_sink(fn: Optional[Callable]) -> None:
|
|
global _byte_sink
|
|
_byte_sink = fn
|
|
|
|
|
|
def register_listener(cb: Listener) -> int:
|
|
"""Register a callback that receives progress events. Returns an id that
|
|
can be passed to `unregister_listener` when the listener is done."""
|
|
with _listener_lock:
|
|
lid = next(_listener_counter)
|
|
_listeners[lid] = cb
|
|
return lid
|
|
|
|
|
|
def unregister_listener(lid: int) -> None:
|
|
with _listener_lock:
|
|
_listeners.pop(lid, None)
|
|
|
|
|
|
def _emit(event: ProgressEvent) -> None:
|
|
"""Fan out to all registered listeners. Never raise — a bad listener
|
|
shouldn't break a download."""
|
|
# Stamp the active repo_id so frontends can route events to the right
|
|
# row. Only set when this emit is happening inside an install handler.
|
|
rid = current_repo_id.get()
|
|
if rid is not None and "repo_id" not in event:
|
|
event = {**event, "repo_id": rid}
|
|
with _listener_lock:
|
|
listeners = list(_listeners.values())
|
|
for cb in listeners:
|
|
try:
|
|
cb(event)
|
|
except Exception as e: # noqa: BLE001
|
|
logger.debug("hf_progress listener raised: %s", e)
|
|
|
|
|
|
def emit(event: ProgressEvent) -> None:
|
|
"""Public emit — lets non-tqdm operations (delete, verify, etc.) push
|
|
lifecycle events onto the same SSE stream."""
|
|
_emit(event)
|
|
|
|
|
|
class SafeFileWrapper:
|
|
def __init__(self, fp):
|
|
self.fp = fp
|
|
self._is_safe_wrapper = True
|
|
def write(self, s):
|
|
try:
|
|
self.fp.write(s)
|
|
except OSError:
|
|
pass
|
|
def flush(self):
|
|
try:
|
|
getattr(self.fp, 'flush', lambda: None)()
|
|
except OSError:
|
|
pass
|
|
def __getattr__(self, name):
|
|
return getattr(self.fp, name)
|
|
|
|
def install() -> None:
|
|
"""Monkey-patch `huggingface_hub`'s tqdm so every download reports to our
|
|
listeners. Safe to call multiple times — second call is a no-op."""
|
|
global _installed
|
|
with _install_lock:
|
|
if _installed:
|
|
return
|
|
# `huggingface_hub.utils.__init__` does `from .tqdm import tqdm`,
|
|
# which shadows the `tqdm` SUBMODULE with the CLASS of the same name
|
|
# when accessed via attribute lookup. Pull the real module out of
|
|
# sys.modules after an explicit import so we patch the right thing.
|
|
try:
|
|
import sys
|
|
import huggingface_hub.utils.tqdm # noqa: F401
|
|
hf_tqdm_module = sys.modules.get("huggingface_hub.utils.tqdm")
|
|
if hf_tqdm_module is None:
|
|
raise ImportError("huggingface_hub.utils.tqdm not in sys.modules after import")
|
|
except Exception as e: # noqa: BLE001
|
|
logger.warning(
|
|
"hf_progress.install: huggingface_hub.utils.tqdm missing (%s); "
|
|
"progress tracking disabled.", e,
|
|
)
|
|
return
|
|
|
|
original = getattr(hf_tqdm_module, "tqdm", None)
|
|
if original is None or not isinstance(original, type):
|
|
logger.warning("hf_progress.install: no `tqdm` class on the module; aborting")
|
|
return
|
|
|
|
class TrackedTqdm(original): # type: ignore[misc,valid-type]
|
|
"""tqdm subclass that emits a progress event on every update."""
|
|
|
|
_last_emit_time: float = 0.0
|
|
|
|
@staticmethod
|
|
def status_printer(file):
|
|
if file is not None and not getattr(file, "_is_safe_wrapper", False):
|
|
file = SafeFileWrapper(file)
|
|
try:
|
|
return original.status_printer(file)
|
|
except Exception:
|
|
return lambda s: None
|
|
|
|
def __init__(self, *args, **kwargs):
|
|
if 'file' in kwargs and kwargs['file'] is not None and not getattr(kwargs['file'], "_is_safe_wrapper", False):
|
|
kwargs['file'] = SafeFileWrapper(kwargs['file'])
|
|
try:
|
|
super().__init__(*args, **kwargs)
|
|
except OSError:
|
|
pass
|
|
|
|
if hasattr(self, 'fp') and getattr(self, 'fp', None) is not None and not getattr(self.fp, "_is_safe_wrapper", False):
|
|
self.fp = SafeFileWrapper(self.fp)
|
|
|
|
import time as _t
|
|
self._last_emit_time = _t.monotonic()
|
|
try:
|
|
desc = getattr(self, "desc", None)
|
|
total = int(getattr(self, "total", 0) or 0)
|
|
_emit({
|
|
"filename": str(desc or "download"),
|
|
"downloaded": 0,
|
|
"total": total,
|
|
"pct": 0.0,
|
|
"phase": "start",
|
|
})
|
|
except Exception:
|
|
pass
|
|
|
|
def _emit_progress(self):
|
|
"""Emit current state as a progress event."""
|
|
try:
|
|
desc = getattr(self, "desc", None)
|
|
total = int(getattr(self, "total", 0) or 0)
|
|
done = int(getattr(self, "n", 0) or 0)
|
|
pct = (done / total) if total > 0 else 0.0
|
|
# Pull rate from tqdm's own calculations if available
|
|
rate = None
|
|
try:
|
|
rate = self.format_dict.get("rate")
|
|
except Exception:
|
|
pass
|
|
event = {
|
|
"filename": str(desc or "download"),
|
|
"downloaded": done,
|
|
"total": total,
|
|
"pct": pct,
|
|
"phase": "done" if (total > 0 and done >= total) else "progress",
|
|
}
|
|
if rate and rate > 0:
|
|
event["rate"] = rate # bytes/sec from tqdm
|
|
_emit(event)
|
|
# Feed the overall aggregator (FDL-06), if wired.
|
|
self._feed_sink(done, total, complete=False)
|
|
except Exception:
|
|
pass
|
|
|
|
def _feed_sink(self, done, total, *, complete: bool):
|
|
"""Forward a byte/count update to the overall aggregator sink.
|
|
|
|
Passes the tqdm `unit` so the aggregator can tell a byte bar
|
|
(unit 'B') from the "Fetching N files" count bar, and a stable
|
|
per-bar key (id(self)) because every per-file byte bar shares
|
|
the default desc 'download' under Xet — keying by desc would
|
|
collapse them into one.
|
|
"""
|
|
sink = _byte_sink
|
|
if sink is None:
|
|
return
|
|
rid = current_repo_id.get()
|
|
if not rid:
|
|
return
|
|
try:
|
|
unit = getattr(self, "unit", None)
|
|
sink(rid, id(self), unit, int(done or 0), int(total or 0), complete)
|
|
except Exception:
|
|
pass
|
|
|
|
def update(self, n=1):
|
|
try:
|
|
super().update(n)
|
|
except OSError:
|
|
pass
|
|
import time as _t
|
|
now = _t.monotonic()
|
|
# Throttle: emit at most every 0.3s to avoid flooding SSE
|
|
if (now - self._last_emit_time) >= 0.3:
|
|
self._last_emit_time = now
|
|
self._emit_progress()
|
|
|
|
def display(self, msg=None, pos=None):
|
|
"""tqdm calls display() on its refresh cycle; piggyback for
|
|
periodic emits even when update() intervals are large."""
|
|
import time as _t
|
|
now = _t.monotonic()
|
|
if (now - self._last_emit_time) >= 0.5:
|
|
self._last_emit_time = now
|
|
self._emit_progress()
|
|
try:
|
|
return super().display(msg, pos)
|
|
except OSError:
|
|
pass
|
|
|
|
def close(self):
|
|
# Credit the file's full size to the aggregator on close. Under
|
|
# Xet a per-file byte bar often never increments `n` (Xet fetches
|
|
# chunks out-of-band), so completion is the only reliable signal
|
|
# that the file's bytes landed. Harmless for classic LFS bars
|
|
# (n already == total).
|
|
try:
|
|
total = int(getattr(self, "total", 0) or 0)
|
|
if total > 0:
|
|
self._feed_sink(total, total, complete=True)
|
|
except Exception:
|
|
pass
|
|
try:
|
|
super().close()
|
|
except OSError:
|
|
pass
|
|
|
|
# Stash the original for inspection / uninstall, then swap.
|
|
hf_tqdm_module._omnivoice_original_tqdm = original # type: ignore[attr-defined]
|
|
hf_tqdm_module.tqdm = TrackedTqdm # type: ignore[assignment]
|
|
global _tracked_tqdm_class
|
|
_tracked_tqdm_class = TrackedTqdm
|
|
_installed = True
|
|
logger.info("hf_progress: installed tqdm patch on huggingface_hub.utils.tqdm")
|