mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 00:55:14 -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"
|
||||||
|
|||||||
+15
-12
@@ -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."""
|
||||||
@@ -93,17 +96,17 @@ def cache_voice_prompt(
|
|||||||
def clear_voice_prompt_cache() -> int:
|
def clear_voice_prompt_cache() -> int:
|
||||||
"""
|
"""
|
||||||
Clear all voice prompt caches (memory and disk).
|
Clear all voice prompt caches (memory and disk).
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Number of cache files deleted
|
Number of cache files deleted
|
||||||
"""
|
"""
|
||||||
# Clear memory cache
|
# Clear memory cache
|
||||||
_memory_cache.clear()
|
_memory_cache.clear()
|
||||||
|
|
||||||
# Clear disk cache
|
# Clear disk cache
|
||||||
cache_dir = _get_cache_dir()
|
cache_dir = _get_cache_dir()
|
||||||
deleted_count = 0
|
deleted_count = 0
|
||||||
|
|
||||||
if cache_dir.exists():
|
if cache_dir.exists():
|
||||||
# Delete prompt cache files
|
# Delete prompt cache files
|
||||||
for cache_file in cache_dir.glob("*.prompt"):
|
for cache_file in cache_dir.glob("*.prompt"):
|
||||||
@@ -111,32 +114,32 @@ 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"):
|
||||||
try:
|
try:
|
||||||
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
|
||||||
|
|
||||||
|
|
||||||
def clear_profile_cache(profile_id: str) -> int:
|
def clear_profile_cache(profile_id: str) -> int:
|
||||||
"""
|
"""
|
||||||
Clear cache files for a specific profile.
|
Clear cache files for a specific profile.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
profile_id: Profile ID
|
profile_id: Profile ID
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Number of cache files deleted
|
Number of cache files deleted
|
||||||
"""
|
"""
|
||||||
cache_dir = _get_cache_dir()
|
cache_dir = _get_cache_dir()
|
||||||
deleted_count = 0
|
deleted_count = 0
|
||||||
|
|
||||||
if cache_dir.exists():
|
if cache_dir.exists():
|
||||||
# Delete combined audio files for this profile
|
# Delete combined audio files for this profile
|
||||||
pattern = f"combined_{profile_id}_*.wav"
|
pattern = f"combined_{profile_id}_*.wav"
|
||||||
@@ -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
|
||||||
|
|||||||
+118
-80
@@ -4,13 +4,16 @@ 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."""
|
||||||
|
|
||||||
def __init__(self, progress_callback: Optional[Callable] = None, filter_non_downloads: bool = False):
|
def __init__(self, progress_callback: Optional[Callable] = None, filter_non_downloads: bool = False):
|
||||||
self.progress_callback = progress_callback
|
self.progress_callback = progress_callback
|
||||||
self.filter_non_downloads = filter_non_downloads # Only filter if True
|
self.filter_non_downloads = filter_non_downloads # Only filter if True
|
||||||
@@ -23,12 +26,12 @@ class HFProgressTracker:
|
|||||||
self._current_filename = ""
|
self._current_filename = ""
|
||||||
self._active_tqdms = {} # Track active tqdm instances
|
self._active_tqdms = {} # Track active tqdm instances
|
||||||
self._hf_tqdm_original_update = None # For monkey-patching hf's tqdm
|
self._hf_tqdm_original_update = None # For monkey-patching hf's tqdm
|
||||||
|
|
||||||
def _create_tracked_tqdm_class(self):
|
def _create_tracked_tqdm_class(self):
|
||||||
"""Create a tqdm subclass that tracks progress."""
|
"""Create a tqdm subclass that tracks progress."""
|
||||||
tracker = self
|
tracker = self
|
||||||
original_tqdm = self._original_tqdm_class
|
original_tqdm = self._original_tqdm_class
|
||||||
|
|
||||||
class TrackedTqdm(original_tqdm):
|
class TrackedTqdm(original_tqdm):
|
||||||
"""A tqdm subclass that reports progress to our tracker."""
|
"""A tqdm subclass that reports progress to our tracker."""
|
||||||
|
|
||||||
@@ -39,7 +42,7 @@ class HFProgressTracker:
|
|||||||
first_arg = args[0]
|
first_arg = args[0]
|
||||||
if isinstance(first_arg, str):
|
if isinstance(first_arg, str):
|
||||||
desc = first_arg
|
desc = first_arg
|
||||||
|
|
||||||
filename = ""
|
filename = ""
|
||||||
if desc:
|
if desc:
|
||||||
# Try to extract filename from description
|
# Try to extract filename from description
|
||||||
@@ -48,44 +51,68 @@ class HFProgressTracker:
|
|||||||
filename = desc.split(":")[0].strip()
|
filename = desc.split(":")[0].strip()
|
||||||
else:
|
else:
|
||||||
filename = desc.strip()
|
filename = desc.strip()
|
||||||
|
|
||||||
# Filter out non-standard kwargs that huggingface_hub might pass
|
# Filter out non-standard kwargs that huggingface_hub might pass
|
||||||
# These are custom kwargs that tqdm doesn't understand
|
# These are custom kwargs that tqdm doesn't understand
|
||||||
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:
|
||||||
filtered_kwargs[key] = value
|
filtered_kwargs[key] = value
|
||||||
|
|
||||||
# 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"
|
||||||
|
|
||||||
with tracker._lock:
|
with tracker._lock:
|
||||||
if filename:
|
if filename:
|
||||||
tracker._current_filename = filename
|
tracker._current_filename = filename
|
||||||
tracker._active_tqdms[id(self)] = {
|
tracker._active_tqdms[id(self)] = {
|
||||||
"filename": self._tracker_filename,
|
"filename": self._tracker_filename,
|
||||||
}
|
}
|
||||||
|
|
||||||
def update(self, n=1):
|
def update(self, n=1):
|
||||||
result = super().update(n)
|
result = super().update(n)
|
||||||
|
|
||||||
@@ -95,95 +122,97 @@ class HFProgressTracker:
|
|||||||
filename = tracker._active_tqdms[id(self)]["filename"]
|
filename = tracker._active_tqdms[id(self)]["filename"]
|
||||||
current = getattr(self, "n", 0)
|
current = getattr(self, "n", 0)
|
||||||
total = getattr(self, "total", 0)
|
total = getattr(self, "total", 0)
|
||||||
|
|
||||||
if total and total > 0:
|
if total and total > 0:
|
||||||
# Always filter out non-byte progress bars (e.g., "Fetching 12 files")
|
# Always filter out non-byte progress bars (e.g., "Fetching 12 files")
|
||||||
# These cause crazy percentages because they're counting files, not bytes
|
# These cause crazy percentages because they're counting files, not bytes
|
||||||
if self._is_non_byte_progress(filename):
|
if self._is_non_byte_progress(filename):
|
||||||
return result
|
return result
|
||||||
|
|
||||||
# When model is cached, also filter out generation-related progress
|
# When model is cached, also filter out generation-related progress
|
||||||
if tracker.filter_non_downloads:
|
if tracker.filter_non_downloads:
|
||||||
if not self._is_download_progress(filename):
|
if not self._is_download_progress(filename):
|
||||||
return result
|
return result
|
||||||
|
|
||||||
# Update per-file tracking
|
# Update per-file tracking
|
||||||
tracker._file_sizes[filename] = total
|
tracker._file_sizes[filename] = total
|
||||||
tracker._file_downloaded[filename] = current
|
tracker._file_downloaded[filename] = current
|
||||||
|
|
||||||
# Calculate totals across all files
|
# Calculate totals across all files
|
||||||
tracker._total_size = sum(tracker._file_sizes.values())
|
tracker._total_size = sum(tracker._file_sizes.values())
|
||||||
tracker._total_downloaded = sum(tracker._file_downloaded.values())
|
tracker._total_downloaded = sum(tracker._file_downloaded.values())
|
||||||
|
|
||||||
# Only report progress once we have a meaningful total (at least 1MB)
|
# Only report progress once we have a meaningful total (at least 1MB)
|
||||||
# This avoids the "100% at 0MB" issue when small config
|
# This avoids the "100% at 0MB" issue when small config
|
||||||
# files are counted before the real model files
|
# files are counted before the real model files
|
||||||
MIN_TOTAL_BYTES = 1_000_000 # 1MB
|
MIN_TOTAL_BYTES = 1_000_000 # 1MB
|
||||||
if tracker._total_size < MIN_TOTAL_BYTES:
|
if tracker._total_size < MIN_TOTAL_BYTES:
|
||||||
return result
|
return result
|
||||||
|
|
||||||
# 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
|
||||||
|
|
||||||
def _is_non_byte_progress(self, filename: str) -> bool:
|
def _is_non_byte_progress(self, filename: str) -> bool:
|
||||||
"""Check if this progress bar should be SKIPPED (returns True to skip).
|
"""Check if this progress bar should be SKIPPED (returns True to skip).
|
||||||
|
|
||||||
We want to track byte-based progress bars. This method identifies
|
We want to track byte-based progress bars. This method identifies
|
||||||
progress bars that count files/items instead of bytes, which would
|
progress bars that count files/items instead of bytes, which would
|
||||||
cause crazy percentages if mixed with our byte counting.
|
cause crazy percentages if mixed with our byte counting.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
True = SKIP this bar (it's not byte-based)
|
True = SKIP this bar (it's not byte-based)
|
||||||
False = TRACK this bar (it counts bytes)
|
False = TRACK this bar (it counts bytes)
|
||||||
"""
|
"""
|
||||||
if not filename:
|
if not filename:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
filename_lower = filename.lower()
|
filename_lower = filename.lower()
|
||||||
|
|
||||||
# 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)
|
||||||
|
|
||||||
def _is_download_progress(self, filename: str) -> bool:
|
def _is_download_progress(self, filename: str) -> bool:
|
||||||
"""Check if this is a real file download progress bar vs internal processing."""
|
"""Check if this is a real file download progress bar vs internal processing."""
|
||||||
if not filename or filename == "unknown":
|
if not filename or filename == "unknown":
|
||||||
return False
|
return False
|
||||||
|
|
||||||
# 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
|
||||||
|
|
||||||
def close(self):
|
def close(self):
|
||||||
with tracker._lock:
|
with tracker._lock:
|
||||||
if id(self) in tracker._active_tqdms:
|
if id(self) in tracker._active_tqdms:
|
||||||
del tracker._active_tqdms[id(self)]
|
del tracker._active_tqdms[id(self)]
|
||||||
return super().close()
|
return super().close()
|
||||||
|
|
||||||
return TrackedTqdm
|
return TrackedTqdm
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def patch_download(self):
|
def patch_download(self):
|
||||||
"""Context manager to patch tqdm for progress tracking."""
|
"""Context manager to patch tqdm for progress tracking."""
|
||||||
@@ -192,7 +221,7 @@ class HFProgressTracker:
|
|||||||
|
|
||||||
# Store original tqdm class
|
# Store original tqdm class
|
||||||
self._original_tqdm_class = tqdm_module.tqdm
|
self._original_tqdm_class = tqdm_module.tqdm
|
||||||
|
|
||||||
# Reset totals
|
# Reset totals
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._total_downloaded = 0
|
self._total_downloaded = 0
|
||||||
@@ -201,7 +230,7 @@ class HFProgressTracker:
|
|||||||
self._file_downloaded = {}
|
self._file_downloaded = {}
|
||||||
self._current_filename = ""
|
self._current_filename = ""
|
||||||
self._active_tqdms = {}
|
self._active_tqdms = {}
|
||||||
|
|
||||||
# Create our tracked tqdm class
|
# Create our tracked tqdm class
|
||||||
tracked_tqdm = self._create_tracked_tqdm_class()
|
tracked_tqdm = self._create_tracked_tqdm_class()
|
||||||
|
|
||||||
@@ -213,13 +242,13 @@ class HFProgressTracker:
|
|||||||
if hasattr(tqdm_module, "auto") and hasattr(tqdm_module.auto, "tqdm"):
|
if hasattr(tqdm_module, "auto") and hasattr(tqdm_module.auto, "tqdm"):
|
||||||
self._original_tqdm_auto = tqdm_module.auto.tqdm
|
self._original_tqdm_auto = tqdm_module.auto.tqdm
|
||||||
tqdm_module.auto.tqdm = tracked_tqdm
|
tqdm_module.auto.tqdm = tracked_tqdm
|
||||||
|
|
||||||
# Patch in sys.modules to catch already-imported references
|
# Patch in sys.modules to catch already-imported references
|
||||||
# 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()):
|
||||||
if "huggingface" in module_name or module_name.startswith("tqdm"):
|
if "huggingface" in module_name or module_name.startswith("tqdm"):
|
||||||
@@ -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}"
|
||||||
@@ -242,31 +274,33 @@ class HFProgressTracker:
|
|||||||
patched_count += 1
|
patched_count += 1
|
||||||
except (AttributeError, TypeError):
|
except (AttributeError, TypeError):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
# ALSO monkey-patch the update method on huggingface_hub's tqdm class
|
# ALSO monkey-patch the update method on huggingface_hub's tqdm class
|
||||||
# This is needed because the class was already defined at import time
|
# This is needed because the class was already defined at import time
|
||||||
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)
|
||||||
# This avoids the "100% at 0MB" issue when small config
|
# This avoids the "100% at 0MB" issue when small config
|
||||||
# files are counted before the real model files
|
# files are counted before the real model files
|
||||||
@@ -274,22 +308,22 @@ class HFProgressTracker:
|
|||||||
if total >= MIN_TOTAL_BYTES:
|
if total >= MIN_TOTAL_BYTES:
|
||||||
tracker._total_downloaded = current
|
tracker._total_downloaded = current
|
||||||
tracker._total_size = total
|
tracker._total_size = total
|
||||||
|
|
||||||
if tracker.progress_callback:
|
if tracker.progress_callback:
|
||||||
tracker.progress_callback(current, total, desc)
|
tracker.progress_callback(current, total, desc)
|
||||||
|
|
||||||
return result
|
return result
|
||||||
|
|
||||||
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
|
||||||
|
|
||||||
except ImportError:
|
except ImportError:
|
||||||
# If tqdm not available, just yield without patching
|
# If tqdm not available, just yield without patching
|
||||||
yield
|
yield
|
||||||
@@ -298,11 +332,12 @@ 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:
|
||||||
tqdm_module.auto.tqdm = self._original_tqdm_auto
|
tqdm_module.auto.tqdm = self._original_tqdm_auto
|
||||||
|
|
||||||
# Restore patched modules
|
# Restore patched modules
|
||||||
for key, (module, attr_name, original) in self._patched_modules.items():
|
for key, (module, attr_name, original) in self._patched_modules.items():
|
||||||
try:
|
try:
|
||||||
@@ -311,26 +346,28 @@ class HFProgressTracker:
|
|||||||
except (AttributeError, TypeError):
|
except (AttributeError, TypeError):
|
||||||
pass
|
pass
|
||||||
self._patched_modules = {}
|
self._patched_modules = {}
|
||||||
|
|
||||||
# Restore hf_tqdm's original update method
|
# Restore hf_tqdm's original update method
|
||||||
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
|
||||||
self._hf_tqdm_original_update = None
|
self._hf_tqdm_original_update = None
|
||||||
|
|
||||||
except (ImportError, AttributeError):
|
except (ImportError, AttributeError):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
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.
|
||||||
|
|
||||||
Note: We send updates even when total=0 (unknown) to provide feedback
|
Note: We send updates even when total=0 (unknown) to provide feedback
|
||||||
during the "incomplete total" phase of huggingface_hub downloads.
|
during the "incomplete total" phase of huggingface_hub downloads.
|
||||||
The frontend handles total=0 gracefully.
|
The frontend handles total=0 gracefully.
|
||||||
@@ -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