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."""
|
"""Pre-download and verify all ML models for the Docker image.
|
||||||
import sys
|
|
||||||
|
|
||||||
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",
|
"u2net",
|
||||||
"isnet-general-use",
|
"isnet-general-use",
|
||||||
"bria-rmbg",
|
"bria-rmbg",
|
||||||
@@ -10,18 +23,99 @@ MODELS = [
|
|||||||
"birefnet-general",
|
"birefnet-general",
|
||||||
]
|
]
|
||||||
|
|
||||||
try:
|
PADDLEOCR_LANGUAGES = ["en", "de", "fr", "es", "zh", "ja", "ko"]
|
||||||
from rembg import new_session
|
|
||||||
except ImportError:
|
|
||||||
print("WARNING: rembg not installed, skipping model pre-download")
|
|
||||||
sys.exit(0)
|
|
||||||
|
|
||||||
for model in MODELS:
|
|
||||||
print(f"Downloading {model}...")
|
def download_rembg_models():
|
||||||
try:
|
"""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)
|
new_session(model)
|
||||||
print(f" {model} ready")
|
print(f" {model} ready")
|
||||||
except Exception as e:
|
print(f"All {len(REMBG_MODELS)} rembg models downloaded.\n")
|
||||||
print(f" WARNING: {model} failed: {e}")
|
|
||||||
|
|
||||||
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