mirror of
https://github.com/snapotter-hq/SnapOtter.git
synced 2026-08-03 07:46:42 +02:00
feat: expand model pre-download with verification and smoke test
Downloads all rembg models (6), RealESRGAN_x4plus.pth weights, PaddleOCR models for all 7 supported languages, verifies MediaPipe bundles its face detection models. Runs a final smoke test importing every ML library. Any failure exits non-zero, failing the Docker build.
This commit is contained in:
+108
-14
@@ -1,7 +1,20 @@
|
||||
"""Pre-download all rembg models offered in the UI."""
|
||||
import sys
|
||||
"""Pre-download and verify all ML models for the Docker image.
|
||||
|
||||
MODELS = [
|
||||
This script runs at Docker build time. Any failure exits non-zero,
|
||||
failing the build. No silent fallbacks.
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
import urllib.request
|
||||
|
||||
REALESRGAN_MODEL_DIR = "/opt/models/realesrgan"
|
||||
REALESRGAN_MODEL_URL = (
|
||||
"https://github.com/xinntao/Real-ESRGAN/releases/download/v0.1.0/RealESRGAN_x4plus.pth"
|
||||
)
|
||||
REALESRGAN_MODEL_PATH = os.path.join(REALESRGAN_MODEL_DIR, "RealESRGAN_x4plus.pth")
|
||||
REALESRGAN_MIN_SIZE = 60_000_000 # ~67 MB
|
||||
|
||||
REMBG_MODELS = [
|
||||
"u2net",
|
||||
"isnet-general-use",
|
||||
"bria-rmbg",
|
||||
@@ -10,18 +23,99 @@ MODELS = [
|
||||
"birefnet-general",
|
||||
]
|
||||
|
||||
try:
|
||||
from rembg import new_session
|
||||
except ImportError:
|
||||
print("WARNING: rembg not installed, skipping model pre-download")
|
||||
sys.exit(0)
|
||||
PADDLEOCR_LANGUAGES = ["en", "de", "fr", "es", "zh", "ja", "ko"]
|
||||
|
||||
for model in MODELS:
|
||||
print(f"Downloading {model}...")
|
||||
try:
|
||||
|
||||
def download_rembg_models():
|
||||
"""Download all rembg ONNX models."""
|
||||
print("=== Downloading rembg models ===")
|
||||
from rembg import new_session
|
||||
|
||||
for model in REMBG_MODELS:
|
||||
print(f" Downloading {model}...")
|
||||
new_session(model)
|
||||
print(f" {model} ready")
|
||||
except Exception as e:
|
||||
print(f" WARNING: {model} failed: {e}")
|
||||
print(f"All {len(REMBG_MODELS)} rembg models downloaded.\n")
|
||||
|
||||
print("Model pre-download complete")
|
||||
|
||||
def download_realesrgan_model():
|
||||
"""Download RealESRGAN_x4plus.pth pretrained weights."""
|
||||
print("=== Downloading RealESRGAN model ===")
|
||||
os.makedirs(REALESRGAN_MODEL_DIR, exist_ok=True)
|
||||
print(f" Downloading from {REALESRGAN_MODEL_URL}...")
|
||||
urllib.request.urlretrieve(REALESRGAN_MODEL_URL, REALESRGAN_MODEL_PATH)
|
||||
|
||||
size = os.path.getsize(REALESRGAN_MODEL_PATH)
|
||||
assert size > REALESRGAN_MIN_SIZE, (
|
||||
f"RealESRGAN model too small: {size} bytes (expected > {REALESRGAN_MIN_SIZE})"
|
||||
)
|
||||
print(f" RealESRGAN_x4plus.pth downloaded ({size / 1_000_000:.1f} MB)\n")
|
||||
|
||||
|
||||
def download_paddleocr_models():
|
||||
"""Pre-download PaddleOCR models for all supported languages."""
|
||||
print("=== Downloading PaddleOCR models ===")
|
||||
os.environ["PADDLE_PDX_DISABLE_MODEL_SOURCE_CHECK"] = "True"
|
||||
from paddleocr import PaddleOCR
|
||||
|
||||
for lang in PADDLEOCR_LANGUAGES:
|
||||
print(f" Downloading models for lang={lang}...")
|
||||
PaddleOCR(lang=lang, use_gpu=False, show_log=False)
|
||||
print(f" {lang} ready")
|
||||
print(f"All {len(PADDLEOCR_LANGUAGES)} PaddleOCR languages downloaded.\n")
|
||||
|
||||
|
||||
def verify_mediapipe():
|
||||
"""Verify MediaPipe face detection models are bundled in the wheel."""
|
||||
print("=== Verifying MediaPipe models ===")
|
||||
import mediapipe as mp
|
||||
|
||||
for selection in [0, 1]:
|
||||
label = "short-range" if selection == 0 else "full-range"
|
||||
print(f" Verifying {label} model (selection={selection})...")
|
||||
detector = mp.solutions.face_detection.FaceDetection(
|
||||
model_selection=selection, min_detection_confidence=0.5
|
||||
)
|
||||
detector.close()
|
||||
print(f" {label} model OK")
|
||||
print("MediaPipe models verified.\n")
|
||||
|
||||
|
||||
def smoke_test():
|
||||
"""Final verification that all ML libraries and models are loadable."""
|
||||
print("=== Running smoke test ===")
|
||||
|
||||
from rembg import new_session
|
||||
from realesrgan import RealESRGANer
|
||||
from basicsr.archs.rrdbnet_arch import RRDBNet
|
||||
from paddleocr import PaddleOCR
|
||||
import mediapipe as mp
|
||||
import cv2
|
||||
import numpy
|
||||
from PIL import Image
|
||||
import seam_carving
|
||||
|
||||
assert os.path.exists(REALESRGAN_MODEL_PATH), (
|
||||
f"RealESRGAN model missing: {REALESRGAN_MODEL_PATH}"
|
||||
)
|
||||
assert os.path.getsize(REALESRGAN_MODEL_PATH) > REALESRGAN_MIN_SIZE, (
|
||||
"RealESRGAN model file is too small"
|
||||
)
|
||||
|
||||
print(" All imports OK")
|
||||
print(" RealESRGAN model file verified")
|
||||
print("Smoke test passed.\n")
|
||||
|
||||
|
||||
def main():
|
||||
print("Pre-downloading all ML models...\n")
|
||||
download_rembg_models()
|
||||
download_realesrgan_model()
|
||||
download_paddleocr_models()
|
||||
verify_mediapipe()
|
||||
smoke_test()
|
||||
print("All models downloaded and verified.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
Reference in New Issue
Block a user