Files
VoiceStudio/backend/utils/hf_progress.py
T
4cc55ab852 Fast model downloads: Xet fast path + accurate progress (FDL W0–W2 + W4) (#424)
* 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>
2026-06-13 20:15:15 +05:30

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")