mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 09:05:17 -07:00
improve startup logging: version, platform, data dir, db stats
Replace verbose startup messages with a clean summary: - App version, Python version, OS/arch - Database path (fix None display), data directory - Profile and generation counts - Backend, GPU, model cache path - Clean up stale loading_model status on startup - Remove noisy progress manager log line
This commit is contained in:
+59
-8
@@ -3,8 +3,36 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
|
import sys
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
class ColoredFormatter(logging.Formatter):
|
||||||
|
"""Custom formatter to add colors matching uvicorn's style."""
|
||||||
|
|
||||||
|
COLORS = {
|
||||||
|
"DEBUG": "\033[36m", # Cyan
|
||||||
|
"INFO": "\033[32m", # Green
|
||||||
|
"WARNING": "\033[33m", # Yellow
|
||||||
|
"ERROR": "\033[31m", # Red
|
||||||
|
"CRITICAL": "\033[35m", # Magenta
|
||||||
|
}
|
||||||
|
RESET = "\033[0m"
|
||||||
|
|
||||||
|
def format(self, record):
|
||||||
|
log_color = self.COLORS.get(record.levelname, self.RESET)
|
||||||
|
record.levelname = f"{log_color}{record.levelname}{self.RESET}"
|
||||||
|
return super().format(record)
|
||||||
|
|
||||||
|
|
||||||
|
# Configure logging to match uvicorn's format with colors
|
||||||
|
handler = logging.StreamHandler(sys.stderr)
|
||||||
|
handler.setFormatter(ColoredFormatter("%(levelname)s: %(message)s"))
|
||||||
|
logging.basicConfig(
|
||||||
|
level=logging.INFO,
|
||||||
|
handlers=[handler],
|
||||||
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
# AMD GPU environment variables must be set before torch import
|
# AMD GPU environment variables must be set before torch import
|
||||||
@@ -18,7 +46,7 @@ from fastapi import FastAPI
|
|||||||
from fastapi.middleware.cors import CORSMiddleware
|
from fastapi.middleware.cors import CORSMiddleware
|
||||||
from urllib.parse import quote
|
from urllib.parse import quote
|
||||||
|
|
||||||
from . import __version__, database
|
from . import __version__, config, database
|
||||||
from .services import tts, transcribe
|
from .services import tts, transcribe
|
||||||
from .database import get_db
|
from .database import get_db
|
||||||
from .utils.platform_detect import get_backend_type
|
from .utils.platform_detect import get_backend_type
|
||||||
@@ -97,9 +125,24 @@ def _register_lifecycle(application: FastAPI) -> None:
|
|||||||
|
|
||||||
@application.on_event("startup")
|
@application.on_event("startup")
|
||||||
async def startup_event():
|
async def startup_event():
|
||||||
logger.info("Voicebox server starting up...")
|
import platform
|
||||||
|
import sys
|
||||||
|
|
||||||
|
logger.info("Voicebox v%s starting up", __version__)
|
||||||
|
logger.info(
|
||||||
|
"Python %s on %s %s (%s)",
|
||||||
|
sys.version.split()[0],
|
||||||
|
platform.system(),
|
||||||
|
platform.release(),
|
||||||
|
platform.machine(),
|
||||||
|
)
|
||||||
|
|
||||||
database.init_db()
|
database.init_db()
|
||||||
logger.info("Database initialized at %s", database._db_path)
|
|
||||||
|
from .database.session import _db_path
|
||||||
|
|
||||||
|
logger.info("Database: %s", _db_path)
|
||||||
|
logger.info("Data directory: %s", config.get_data_dir())
|
||||||
|
|
||||||
init_queue()
|
init_queue()
|
||||||
|
|
||||||
@@ -112,11 +155,19 @@ def _register_lifecycle(application: FastAPI) -> None:
|
|||||||
sa_text(
|
sa_text(
|
||||||
"UPDATE generations SET status = 'failed', "
|
"UPDATE generations SET status = 'failed', "
|
||||||
"error = 'Server was shut down during generation' "
|
"error = 'Server was shut down during generation' "
|
||||||
"WHERE status = 'generating'"
|
"WHERE status IN ('generating', 'loading_model')"
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
if result.rowcount > 0:
|
if result.rowcount > 0:
|
||||||
logger.info("Marked %d stale generation(s) as failed", result.rowcount)
|
logger.info("Marked %d stale generation(s) as failed", result.rowcount)
|
||||||
|
|
||||||
|
# Log database stats
|
||||||
|
from .database import VoiceProfile as DBVoiceProfile, Generation as DBGeneration
|
||||||
|
|
||||||
|
profile_count = db.query(DBVoiceProfile).count()
|
||||||
|
generation_count = db.query(DBGeneration).count()
|
||||||
|
logger.info("Profiles: %d, Generations: %d", profile_count, generation_count)
|
||||||
|
|
||||||
db.commit()
|
db.commit()
|
||||||
db.close()
|
db.close()
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -124,7 +175,7 @@ def _register_lifecycle(application: FastAPI) -> None:
|
|||||||
|
|
||||||
backend_type = get_backend_type()
|
backend_type = get_backend_type()
|
||||||
logger.info("Backend: %s", backend_type.upper())
|
logger.info("Backend: %s", backend_type.upper())
|
||||||
logger.info("GPU available: %s", _get_gpu_status())
|
logger.info("GPU: %s", _get_gpu_status())
|
||||||
|
|
||||||
from .services.cuda import check_and_update_cuda_binary
|
from .services.cuda import check_and_update_cuda_binary
|
||||||
|
|
||||||
@@ -133,7 +184,6 @@ def _register_lifecycle(application: FastAPI) -> None:
|
|||||||
try:
|
try:
|
||||||
progress_manager = get_progress_manager()
|
progress_manager = get_progress_manager()
|
||||||
progress_manager._set_main_loop(asyncio.get_running_loop())
|
progress_manager._set_main_loop(asyncio.get_running_loop())
|
||||||
logger.info("Progress manager initialized with event loop")
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning("Could not initialize progress manager event loop: %s", e)
|
logger.warning("Could not initialize progress manager event loop: %s", e)
|
||||||
|
|
||||||
@@ -142,10 +192,11 @@ def _register_lifecycle(application: FastAPI) -> None:
|
|||||||
|
|
||||||
cache_dir = Path(hf_constants.HF_HUB_CACHE)
|
cache_dir = Path(hf_constants.HF_HUB_CACHE)
|
||||||
cache_dir.mkdir(parents=True, exist_ok=True)
|
cache_dir.mkdir(parents=True, exist_ok=True)
|
||||||
logger.info("HuggingFace cache directory: %s", cache_dir)
|
logger.info("Model cache: %s", cache_dir)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning("Could not create HuggingFace cache directory: %s", e)
|
logger.warning("Could not create HuggingFace cache directory: %s", e)
|
||||||
logger.warning("Model downloads may fail. Please ensure the directory exists and has write permissions.")
|
|
||||||
|
logger.info("Ready")
|
||||||
|
|
||||||
@application.on_event("shutdown")
|
@application.on_event("shutdown")
|
||||||
async def shutdown_event():
|
async def shutdown_event():
|
||||||
|
|||||||
@@ -50,7 +50,7 @@ def build_server(cuda=False):
|
|||||||
qwen_tts_path = os.getenv("QWEN_TTS_PATH")
|
qwen_tts_path = os.getenv("QWEN_TTS_PATH")
|
||||||
if qwen_tts_path and Path(qwen_tts_path).exists():
|
if qwen_tts_path and Path(qwen_tts_path).exists():
|
||||||
args.extend(["--paths", str(qwen_tts_path)])
|
args.extend(["--paths", str(qwen_tts_path)])
|
||||||
print(f"Using local qwen_tts source from: {qwen_tts_path}")
|
logger.info("Using local qwen_tts source from: %s", qwen_tts_path)
|
||||||
|
|
||||||
# Add common hidden imports
|
# Add common hidden imports
|
||||||
args.extend(
|
args.extend(
|
||||||
@@ -185,7 +185,7 @@ def build_server(cuda=False):
|
|||||||
|
|
||||||
# Add CUDA-specific hidden imports
|
# Add CUDA-specific hidden imports
|
||||||
if cuda:
|
if cuda:
|
||||||
print("Building with CUDA support")
|
logger.info("Building with CUDA support")
|
||||||
args.extend(
|
args.extend(
|
||||||
[
|
[
|
||||||
"--hidden-import",
|
"--hidden-import",
|
||||||
@@ -219,7 +219,7 @@ def build_server(cuda=False):
|
|||||||
|
|
||||||
# Add MLX-specific imports if building on Apple Silicon (never for CUDA builds)
|
# Add MLX-specific imports if building on Apple Silicon (never for CUDA builds)
|
||||||
if is_apple_silicon() and not cuda:
|
if is_apple_silicon() and not cuda:
|
||||||
print("Building for Apple Silicon - including MLX dependencies")
|
logger.info("Building for Apple Silicon - including MLX dependencies")
|
||||||
args.extend(
|
args.extend(
|
||||||
[
|
[
|
||||||
"--hidden-import",
|
"--hidden-import",
|
||||||
@@ -252,7 +252,7 @@ def build_server(cuda=False):
|
|||||||
]
|
]
|
||||||
)
|
)
|
||||||
elif not cuda:
|
elif not cuda:
|
||||||
print("Building for non-Apple Silicon platform - PyTorch only")
|
logger.info("Building for non-Apple Silicon platform - PyTorch only")
|
||||||
|
|
||||||
dist_dir = str(backend_dir / "dist")
|
dist_dir = str(backend_dir / "dist")
|
||||||
build_dir = str(backend_dir / "build")
|
build_dir = str(backend_dir / "build")
|
||||||
@@ -284,7 +284,7 @@ def build_server(cuda=False):
|
|||||||
)
|
)
|
||||||
has_cuda_torch = bool(result.stdout.strip())
|
has_cuda_torch = bool(result.stdout.strip())
|
||||||
if has_cuda_torch:
|
if has_cuda_torch:
|
||||||
print("CUDA torch detected — installing CPU torch for CPU build...")
|
logger.info("CUDA torch detected — installing CPU torch for CPU build...")
|
||||||
subprocess.run(
|
subprocess.run(
|
||||||
[
|
[
|
||||||
sys.executable,
|
sys.executable,
|
||||||
@@ -309,7 +309,7 @@ def build_server(cuda=False):
|
|||||||
finally:
|
finally:
|
||||||
# Restore CUDA torch if we swapped it out (even on build failure)
|
# Restore CUDA torch if we swapped it out (even on build failure)
|
||||||
if restore_cuda:
|
if restore_cuda:
|
||||||
print("Restoring CUDA torch...")
|
logger.info("Restoring CUDA torch...")
|
||||||
import subprocess
|
import subprocess
|
||||||
|
|
||||||
subprocess.run(
|
subprocess.run(
|
||||||
@@ -329,7 +329,7 @@ def build_server(cuda=False):
|
|||||||
check=True,
|
check=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
print(f"Binary built in {backend_dir / 'dist' / binary_name}")
|
logger.info("Binary built in %s", backend_dir / "dist" / binary_name)
|
||||||
|
|
||||||
|
|
||||||
def _get_cuda_dll_excludes():
|
def _get_cuda_dll_excludes():
|
||||||
@@ -372,7 +372,7 @@ def _get_cuda_dll_excludes():
|
|||||||
total_mb = (
|
total_mb = (
|
||||||
sum((torch_lib / dll).stat().st_size for dll in exclude_dlls if (torch_lib / dll).exists()) / 1024 / 1024
|
sum((torch_lib / dll).stat().st_size for dll in exclude_dlls if (torch_lib / dll).exists()) / 1024 / 1024
|
||||||
)
|
)
|
||||||
print(f"CPU build: will exclude {len(exclude_dlls)} CUDA DLLs ({total_mb:.0f} MB)")
|
logger.info("CPU build: will exclude %d CUDA DLLs (%.0f MB)", len(exclude_dlls), total_mb)
|
||||||
|
|
||||||
return exclude_dlls
|
return exclude_dlls
|
||||||
|
|
||||||
|
|||||||
+12
-2
@@ -4,20 +4,24 @@ Configuration module for voicebox backend.
|
|||||||
Handles data directory configuration for production bundling.
|
Handles data directory configuration for production bundling.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
import os
|
import os
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
# Allow users to override the HuggingFace model download directory.
|
# Allow users to override the HuggingFace model download directory.
|
||||||
# Set VOICEBOX_MODELS_DIR to an absolute path before starting the server.
|
# Set VOICEBOX_MODELS_DIR to an absolute path before starting the server.
|
||||||
# This sets HF_HUB_CACHE so all huggingface_hub downloads go to that path.
|
# This sets HF_HUB_CACHE so all huggingface_hub downloads go to that path.
|
||||||
_custom_models_dir = os.environ.get("VOICEBOX_MODELS_DIR")
|
_custom_models_dir = os.environ.get("VOICEBOX_MODELS_DIR")
|
||||||
if _custom_models_dir:
|
if _custom_models_dir:
|
||||||
os.environ["HF_HUB_CACHE"] = _custom_models_dir
|
os.environ["HF_HUB_CACHE"] = _custom_models_dir
|
||||||
print(f"[config] Model download path set to: {_custom_models_dir}")
|
logger.info("Model download path set to: %s", _custom_models_dir)
|
||||||
|
|
||||||
# Default data directory (used in development)
|
# Default data directory (used in development)
|
||||||
_data_dir = Path("data")
|
_data_dir = Path("data")
|
||||||
|
|
||||||
|
|
||||||
def set_data_dir(path: str | Path):
|
def set_data_dir(path: str | Path):
|
||||||
"""
|
"""
|
||||||
Set the data directory path.
|
Set the data directory path.
|
||||||
@@ -28,7 +32,8 @@ def set_data_dir(path: str | Path):
|
|||||||
global _data_dir
|
global _data_dir
|
||||||
_data_dir = Path(path)
|
_data_dir = Path(path)
|
||||||
_data_dir.mkdir(parents=True, exist_ok=True)
|
_data_dir.mkdir(parents=True, exist_ok=True)
|
||||||
print(f"Data directory set to: {_data_dir.absolute()}")
|
logger.info("Data directory set to: %s", _data_dir.absolute())
|
||||||
|
|
||||||
|
|
||||||
def get_data_dir() -> Path:
|
def get_data_dir() -> Path:
|
||||||
"""
|
"""
|
||||||
@@ -39,28 +44,33 @@ def get_data_dir() -> Path:
|
|||||||
"""
|
"""
|
||||||
return _data_dir
|
return _data_dir
|
||||||
|
|
||||||
|
|
||||||
def get_db_path() -> Path:
|
def get_db_path() -> Path:
|
||||||
"""Get database file path."""
|
"""Get database file path."""
|
||||||
return _data_dir / "voicebox.db"
|
return _data_dir / "voicebox.db"
|
||||||
|
|
||||||
|
|
||||||
def get_profiles_dir() -> Path:
|
def get_profiles_dir() -> Path:
|
||||||
"""Get profiles directory path."""
|
"""Get profiles directory path."""
|
||||||
path = _data_dir / "profiles"
|
path = _data_dir / "profiles"
|
||||||
path.mkdir(parents=True, exist_ok=True)
|
path.mkdir(parents=True, exist_ok=True)
|
||||||
return path
|
return path
|
||||||
|
|
||||||
|
|
||||||
def get_generations_dir() -> Path:
|
def get_generations_dir() -> Path:
|
||||||
"""Get generations directory path."""
|
"""Get generations directory path."""
|
||||||
path = _data_dir / "generations"
|
path = _data_dir / "generations"
|
||||||
path.mkdir(parents=True, exist_ok=True)
|
path.mkdir(parents=True, exist_ok=True)
|
||||||
return path
|
return path
|
||||||
|
|
||||||
|
|
||||||
def get_cache_dir() -> Path:
|
def get_cache_dir() -> Path:
|
||||||
"""Get cache directory path."""
|
"""Get cache directory path."""
|
||||||
path = _data_dir / "cache"
|
path = _data_dir / "cache"
|
||||||
path.mkdir(parents=True, exist_ok=True)
|
path.mkdir(parents=True, exist_ok=True)
|
||||||
return path
|
return path
|
||||||
|
|
||||||
|
|
||||||
def get_models_dir() -> Path:
|
def get_models_dir() -> Path:
|
||||||
"""Get models directory path."""
|
"""Get models directory path."""
|
||||||
path = _data_dir / "models"
|
path = _data_dir / "models"
|
||||||
|
|||||||
@@ -3,12 +3,15 @@ Voice prompt caching utilities.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import hashlib
|
import hashlib
|
||||||
|
import logging
|
||||||
import torch
|
import torch
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Optional, Union, Dict, Any
|
from typing import Optional, Union, Dict, Any
|
||||||
|
|
||||||
from .. import config
|
from .. import config
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
def _get_cache_dir() -> Path:
|
def _get_cache_dir() -> Path:
|
||||||
"""Get cache directory from config."""
|
"""Get cache directory from config."""
|
||||||
@@ -111,7 +114,7 @@ def clear_voice_prompt_cache() -> int:
|
|||||||
cache_file.unlink()
|
cache_file.unlink()
|
||||||
deleted_count += 1
|
deleted_count += 1
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f"Failed to delete cache file {cache_file}: {e}")
|
logger.warning("Failed to delete cache file %s: %s", cache_file, e)
|
||||||
|
|
||||||
# Delete combined audio files
|
# Delete combined audio files
|
||||||
for audio_file in cache_dir.glob("combined_*.wav"):
|
for audio_file in cache_dir.glob("combined_*.wav"):
|
||||||
@@ -119,7 +122,7 @@ def clear_voice_prompt_cache() -> int:
|
|||||||
audio_file.unlink()
|
audio_file.unlink()
|
||||||
deleted_count += 1
|
deleted_count += 1
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f"Failed to delete combined audio file {audio_file}: {e}")
|
logger.warning("Failed to delete combined audio file %s: %s", audio_file, e)
|
||||||
|
|
||||||
return deleted_count
|
return deleted_count
|
||||||
|
|
||||||
@@ -145,6 +148,6 @@ def clear_profile_cache(profile_id: str) -> int:
|
|||||||
audio_file.unlink()
|
audio_file.unlink()
|
||||||
deleted_count += 1
|
deleted_count += 1
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f"Failed to delete combined audio file {audio_file}: {e}")
|
logger.warning("Failed to delete combined audio file %s: %s", audio_file, e)
|
||||||
|
|
||||||
return deleted_count
|
return deleted_count
|
||||||
|
|||||||
@@ -4,9 +4,12 @@ HuggingFace Hub download progress tracking.
|
|||||||
|
|
||||||
from typing import Optional, Callable
|
from typing import Optional, Callable
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
|
import logging
|
||||||
import threading
|
import threading
|
||||||
import sys
|
import sys
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class HFProgressTracker:
|
class HFProgressTracker:
|
||||||
"""Tracks HuggingFace Hub download progress by intercepting tqdm."""
|
"""Tracks HuggingFace Hub download progress by intercepting tqdm."""
|
||||||
@@ -54,11 +57,35 @@ class HFProgressTracker:
|
|||||||
filtered_kwargs = {}
|
filtered_kwargs = {}
|
||||||
# Known tqdm kwargs - pass these through
|
# Known tqdm kwargs - pass these through
|
||||||
tqdm_kwargs = {
|
tqdm_kwargs = {
|
||||||
'iterable', 'desc', 'total', 'leave', 'file', 'ncols', 'mininterval',
|
"iterable",
|
||||||
'maxinterval', 'miniters', 'ascii', 'disable', 'unit', 'unit_scale',
|
"desc",
|
||||||
'dynamic_ncols', 'smoothing', 'bar_format', 'initial', 'position',
|
"total",
|
||||||
'postfix', 'unit_divisor', 'write_bytes', 'lock_args', 'nrows',
|
"leave",
|
||||||
'colour', 'color', 'delay', 'gui', 'disable_default', 'pos'
|
"file",
|
||||||
|
"ncols",
|
||||||
|
"mininterval",
|
||||||
|
"maxinterval",
|
||||||
|
"miniters",
|
||||||
|
"ascii",
|
||||||
|
"disable",
|
||||||
|
"unit",
|
||||||
|
"unit_scale",
|
||||||
|
"dynamic_ncols",
|
||||||
|
"smoothing",
|
||||||
|
"bar_format",
|
||||||
|
"initial",
|
||||||
|
"position",
|
||||||
|
"postfix",
|
||||||
|
"unit_divisor",
|
||||||
|
"write_bytes",
|
||||||
|
"lock_args",
|
||||||
|
"nrows",
|
||||||
|
"colour",
|
||||||
|
"color",
|
||||||
|
"delay",
|
||||||
|
"gui",
|
||||||
|
"disable_default",
|
||||||
|
"pos",
|
||||||
}
|
}
|
||||||
for key, value in kwargs.items():
|
for key, value in kwargs.items():
|
||||||
if key in tqdm_kwargs:
|
if key in tqdm_kwargs:
|
||||||
@@ -67,14 +94,14 @@ class HFProgressTracker:
|
|||||||
# Force-enable the progress bar — we're tracking progress ourselves,
|
# Force-enable the progress bar — we're tracking progress ourselves,
|
||||||
# we don't need tqdm to render to a terminal, but we DO need
|
# we don't need tqdm to render to a terminal, but we DO need
|
||||||
# self.n to be updated when update() is called.
|
# self.n to be updated when update() is called.
|
||||||
filtered_kwargs['disable'] = False
|
filtered_kwargs["disable"] = False
|
||||||
|
|
||||||
# Try to initialize with filtered kwargs, fall back to all kwargs if that fails
|
# Try to initialize with filtered kwargs, fall back to all kwargs if that fails
|
||||||
try:
|
try:
|
||||||
super().__init__(*args, **filtered_kwargs)
|
super().__init__(*args, **filtered_kwargs)
|
||||||
except TypeError:
|
except TypeError:
|
||||||
# If filtering failed, try with all kwargs (maybe tqdm version accepts them)
|
# If filtering failed, try with all kwargs (maybe tqdm version accepts them)
|
||||||
kwargs['disable'] = False
|
kwargs["disable"] = False
|
||||||
super().__init__(*args, **kwargs)
|
super().__init__(*args, **kwargs)
|
||||||
|
|
||||||
self._tracker_filename = filename or "unknown"
|
self._tracker_filename = filename or "unknown"
|
||||||
@@ -124,11 +151,7 @@ class HFProgressTracker:
|
|||||||
|
|
||||||
# Call progress callback
|
# Call progress callback
|
||||||
if tracker.progress_callback:
|
if tracker.progress_callback:
|
||||||
tracker.progress_callback(
|
tracker.progress_callback(tracker._total_downloaded, tracker._total_size, filename)
|
||||||
tracker._total_downloaded,
|
|
||||||
tracker._total_size,
|
|
||||||
filename
|
|
||||||
)
|
|
||||||
|
|
||||||
return result
|
return result
|
||||||
|
|
||||||
@@ -151,7 +174,7 @@ class HFProgressTracker:
|
|||||||
# Skip "Fetching X files" - it counts files (total=12), not bytes
|
# Skip "Fetching X files" - it counts files (total=12), not bytes
|
||||||
# Don't skip "Downloading (incomplete total...)" - that IS byte-based
|
# Don't skip "Downloading (incomplete total...)" - that IS byte-based
|
||||||
skip_patterns = [
|
skip_patterns = [
|
||||||
'fetching', # "Fetching 12 files" has total=12 files, not bytes
|
"fetching", # "Fetching 12 files" has total=12 files, not bytes
|
||||||
]
|
]
|
||||||
return any(pattern in filename_lower for pattern in skip_patterns)
|
return any(pattern in filename_lower for pattern in skip_patterns)
|
||||||
|
|
||||||
@@ -162,16 +185,22 @@ class HFProgressTracker:
|
|||||||
|
|
||||||
# Real downloads have file extensions
|
# Real downloads have file extensions
|
||||||
download_extensions = [
|
download_extensions = [
|
||||||
'.safetensors', '.bin', '.pt', '.pth', # Model weights
|
".safetensors",
|
||||||
'.json', '.txt', '.py', # Config files
|
".bin",
|
||||||
'.msgpack', '.h5', # Other formats
|
".pt",
|
||||||
|
".pth", # Model weights
|
||||||
|
".json",
|
||||||
|
".txt",
|
||||||
|
".py", # Config files
|
||||||
|
".msgpack",
|
||||||
|
".h5", # Other formats
|
||||||
]
|
]
|
||||||
|
|
||||||
filename_lower = filename.lower()
|
filename_lower = filename.lower()
|
||||||
has_extension = any(filename_lower.endswith(ext) for ext in download_extensions)
|
has_extension = any(filename_lower.endswith(ext) for ext in download_extensions)
|
||||||
|
|
||||||
# Skip generation-related progress indicators
|
# Skip generation-related progress indicators
|
||||||
skip_patterns = ['segment', 'processing', 'generating', 'loading']
|
skip_patterns = ["segment", "processing", "generating", "loading"]
|
||||||
has_skip_pattern = any(pattern in filename_lower for pattern in skip_patterns)
|
has_skip_pattern = any(pattern in filename_lower for pattern in skip_patterns)
|
||||||
|
|
||||||
return has_extension and not has_skip_pattern
|
return has_extension and not has_skip_pattern
|
||||||
@@ -218,7 +247,7 @@ class HFProgressTracker:
|
|||||||
# huggingface_hub uses: from tqdm.auto import tqdm as base_tqdm
|
# huggingface_hub uses: from tqdm.auto import tqdm as base_tqdm
|
||||||
# So we need to patch both 'tqdm' and 'base_tqdm' attributes
|
# So we need to patch both 'tqdm' and 'base_tqdm' attributes
|
||||||
self._patched_modules = {}
|
self._patched_modules = {}
|
||||||
tqdm_attr_names = ['tqdm', 'base_tqdm', 'old_tqdm'] # Various names used
|
tqdm_attr_names = ["tqdm", "base_tqdm", "old_tqdm"] # Various names used
|
||||||
|
|
||||||
patched_count = 0
|
patched_count = 0
|
||||||
for module_name in list(sys.modules.keys()):
|
for module_name in list(sys.modules.keys()):
|
||||||
@@ -230,10 +259,13 @@ class HFProgressTracker:
|
|||||||
attr = getattr(module, attr_name)
|
attr = getattr(module, attr_name)
|
||||||
# Only patch if it's a tqdm class (not already patched)
|
# Only patch if it's a tqdm class (not already patched)
|
||||||
is_tqdm_class = (
|
is_tqdm_class = (
|
||||||
attr is self._original_tqdm_class or
|
attr is self._original_tqdm_class
|
||||||
(self._original_tqdm_auto and attr is self._original_tqdm_auto) or
|
or (self._original_tqdm_auto and attr is self._original_tqdm_auto)
|
||||||
(hasattr(attr, "__name__") and attr.__name__ == "tqdm" and
|
or (
|
||||||
hasattr(attr, "update")) # tqdm classes have update method
|
hasattr(attr, "__name__")
|
||||||
|
and attr.__name__ == "tqdm"
|
||||||
|
and hasattr(attr, "update")
|
||||||
|
) # tqdm classes have update method
|
||||||
)
|
)
|
||||||
if is_tqdm_class:
|
if is_tqdm_class:
|
||||||
key = f"{module_name}.{attr_name}"
|
key = f"{module_name}.{attr_name}"
|
||||||
@@ -248,23 +280,25 @@ class HFProgressTracker:
|
|||||||
self._hf_tqdm_original_update = None
|
self._hf_tqdm_original_update = None
|
||||||
try:
|
try:
|
||||||
from huggingface_hub.utils import tqdm as hf_tqdm_module
|
from huggingface_hub.utils import tqdm as hf_tqdm_module
|
||||||
if hasattr(hf_tqdm_module, 'tqdm'):
|
|
||||||
|
if hasattr(hf_tqdm_module, "tqdm"):
|
||||||
hf_tqdm_class = hf_tqdm_module.tqdm
|
hf_tqdm_class = hf_tqdm_module.tqdm
|
||||||
self._hf_tqdm_original_update = hf_tqdm_class.update
|
self._hf_tqdm_original_update = hf_tqdm_class.update
|
||||||
|
|
||||||
# Create a wrapper that calls our tracking
|
# Create a wrapper that calls our tracking
|
||||||
tracker = self # Reference to HFProgressTracker instance
|
tracker = self # Reference to HFProgressTracker instance
|
||||||
|
|
||||||
def patched_update(tqdm_self, n=1):
|
def patched_update(tqdm_self, n=1):
|
||||||
result = tracker._hf_tqdm_original_update(tqdm_self, n)
|
result = tracker._hf_tqdm_original_update(tqdm_self, n)
|
||||||
|
|
||||||
# Track this progress
|
# Track this progress
|
||||||
with tracker._lock:
|
with tracker._lock:
|
||||||
desc = getattr(tqdm_self, 'desc', '') or ''
|
desc = getattr(tqdm_self, "desc", "") or ""
|
||||||
current = getattr(tqdm_self, 'n', 0)
|
current = getattr(tqdm_self, "n", 0)
|
||||||
total = getattr(tqdm_self, 'total', 0) or 0
|
total = getattr(tqdm_self, "total", 0) or 0
|
||||||
|
|
||||||
# Skip non-byte progress bars
|
# Skip non-byte progress bars
|
||||||
if 'fetching' in desc.lower():
|
if "fetching" in desc.lower():
|
||||||
return result
|
return result
|
||||||
|
|
||||||
# Skip until we have a meaningful total (at least 1MB)
|
# Skip until we have a meaningful total (at least 1MB)
|
||||||
@@ -282,11 +316,11 @@ class HFProgressTracker:
|
|||||||
|
|
||||||
hf_tqdm_class.update = patched_update
|
hf_tqdm_class.update = patched_update
|
||||||
patched_count += 1
|
patched_count += 1
|
||||||
print(f"[HFProgressTracker] Monkey-patched huggingface_hub.utils.tqdm.tqdm.update")
|
logger.debug("Monkey-patched huggingface_hub.utils.tqdm.tqdm.update")
|
||||||
except (ImportError, AttributeError) as e:
|
except (ImportError, AttributeError) as e:
|
||||||
print(f"[HFProgressTracker] Could not monkey-patch hf_tqdm: {e}")
|
logger.warning("Could not monkey-patch hf_tqdm: %s", e)
|
||||||
|
|
||||||
print(f"[HFProgressTracker] Patched {patched_count} tqdm references")
|
logger.debug("Patched %d tqdm references", patched_count)
|
||||||
|
|
||||||
yield
|
yield
|
||||||
|
|
||||||
@@ -298,6 +332,7 @@ class HFProgressTracker:
|
|||||||
if self._original_tqdm_class:
|
if self._original_tqdm_class:
|
||||||
try:
|
try:
|
||||||
import tqdm as tqdm_module
|
import tqdm as tqdm_module
|
||||||
|
|
||||||
tqdm_module.tqdm = self._original_tqdm_class
|
tqdm_module.tqdm = self._original_tqdm_class
|
||||||
|
|
||||||
if self._original_tqdm_auto:
|
if self._original_tqdm_auto:
|
||||||
@@ -316,7 +351,8 @@ class HFProgressTracker:
|
|||||||
if self._hf_tqdm_original_update:
|
if self._hf_tqdm_original_update:
|
||||||
try:
|
try:
|
||||||
from huggingface_hub.utils import tqdm as hf_tqdm_module
|
from huggingface_hub.utils import tqdm as hf_tqdm_module
|
||||||
if hasattr(hf_tqdm_module, 'tqdm'):
|
|
||||||
|
if hasattr(hf_tqdm_module, "tqdm"):
|
||||||
hf_tqdm_module.tqdm.update = self._hf_tqdm_original_update
|
hf_tqdm_module.tqdm.update = self._hf_tqdm_original_update
|
||||||
except (ImportError, AttributeError):
|
except (ImportError, AttributeError):
|
||||||
pass
|
pass
|
||||||
@@ -328,6 +364,7 @@ class HFProgressTracker:
|
|||||||
|
|
||||||
def create_hf_progress_callback(model_name: str, progress_manager):
|
def create_hf_progress_callback(model_name: str, progress_manager):
|
||||||
"""Create a progress callback for HuggingFace downloads."""
|
"""Create a progress callback for HuggingFace downloads."""
|
||||||
|
|
||||||
def callback(downloaded: int, total: int, filename: str = ""):
|
def callback(downloaded: int, total: int, filename: str = ""):
|
||||||
"""Progress callback.
|
"""Progress callback.
|
||||||
|
|
||||||
@@ -342,4 +379,5 @@ def create_hf_progress_callback(model_name: str, progress_manager):
|
|||||||
filename=filename or "",
|
filename=filename or "",
|
||||||
status="downloading",
|
status="downloading",
|
||||||
)
|
)
|
||||||
|
|
||||||
return callback
|
return callback
|
||||||
|
|||||||
Reference in New Issue
Block a user