unkor-backend / model_bootstrap.py
wachekk's picture
feat: lazy loading Naman712
492b161
Raw
History Blame Contribute Delete
10.5 kB
"""
model_bootstrap.py — Provisionnement automatique des modèles au démarrage.
Si les modèles ONNX ne sont pas présents localement, ils sont récupérés depuis
HuggingFace (via `huggingface_hub.hf_hub_download`, authentifié par `HF_TOKEN`)
puis préparés, **en arrière-plan** (thread démon) pour ne pas bloquer le serveur.
1. Détecteur IA global -> models/ai_detector.onnx (+ sidecar ai_detector.json)
Source : SDXL_MODEL_REPO (défaut `umm-maybe/AI-image-detector`).
- Si SDXL_MODEL_FILE est défini : téléchargement direct de l'ONNX publié.
- Sinon : téléchargement de pytorch_model.bin + config + preprocessor, puis
conversion PyTorch -> ONNX (transformers/torch) et écriture du sidecar.
2. Modèle visage -> models/deepfake_efficientnet.onnx
Source : ONNX_MODEL_REPO + ONNX_MODEL_FILE (repo HF + fichier .onnx).
Le modèle VIDÉO (ResNext50+LSTM, Naman712) n'est PLUS provisionné ici : il est
téléchargé ET chargé PARESSEUSEMENT par video_model.VideoDeepfakeAnalyzer, au
premier appel à predict() (première analyse vidéo). Voir video_model.py pour
VIDEO_MODEL_REPO / VIDEO_MODEL_FILE / VIDEO_MODEL_PATH / VIDEO_FAKE_INDEX.
Variables d'environnement :
HF_TOKEN jeton HuggingFace (repos privés / gated / quotas)
SDXL_MODEL_PATH chemin .onnx local explicite (court-circuite le DL)
SDXL_MODEL_REPO repo HF du détecteur IA (défaut umm-maybe/AI-image-detector)
SDXL_MODEL_FILE nom d'un .onnx déjà publié dans le repo (téléchargement direct)
ONNX_MODEL_PATH chemin .onnx visage local explicite
ONNX_MODEL_REPO repo HF du modèle visage
ONNX_MODEL_FILE fichier .onnx visage dans le repo
"""
import os
import json
import shutil
import threading
_BASE = os.path.dirname(os.path.abspath(__file__))
_MODELS = os.path.join(_BASE, "models")
AI_ONNX = os.path.join(_MODELS, "ai_detector.onnx")
AI_SIDECAR = os.path.join(_MODELS, "ai_detector.json")
FACE_ONNX = os.path.join(_MODELS, "deepfake_efficientnet.onnx")
YUNET_ONNX = os.path.join(_MODELS, "face_detector_yunet.onnx")
YUNET_URL = os.getenv(
"FACE_DETECTOR_URL",
"https://huggingface.co/opencv/face_detection_yunet/resolve/main/face_detection_yunet_2023mar.onnx")
_DEF_MEAN = [0.485, 0.456, 0.406]
_DEF_STD = [0.229, 0.224, 0.225]
def _log(msg):
print(f"[bootstrap] {msg}", flush=True)
def _hf_token():
return os.getenv("HF_TOKEN") or os.getenv("HUGGINGFACE_TOKEN") or None
# ── 1. Détecteur IA global ──────────────────────────────────────────────────
def ensure_sdxl_model():
if os.path.exists(AI_ONNX):
_log("ai_detector.onnx déjà présent — rien à télécharger.")
return
explicit = os.getenv("SDXL_MODEL_PATH")
if explicit and os.path.exists(explicit):
_log(f"SDXL_MODEL_PATH fourni ({explicit}) — pas de téléchargement.")
return
from huggingface_hub import hf_hub_download
repo = os.getenv("SDXL_MODEL_REPO", "umm-maybe/AI-image-detector")
token = _hf_token()
os.makedirs(_MODELS, exist_ok=True)
# (a) ONNX déjà publié dans le repo -> téléchargement direct
onnx_file = os.getenv("SDXL_MODEL_FILE")
if onnx_file:
try:
f = hf_hub_download(repo_id=repo, filename=onnx_file, token=token)
shutil.copy(f, AI_ONNX)
try:
s = hf_hub_download(repo_id=repo, filename="ai_detector.json", token=token)
shutil.copy(s, AI_SIDECAR)
except Exception:
pass
_log(f"ai_detector.onnx téléchargé directement depuis {repo}/{onnx_file}")
return
except Exception as e:
_log(f"ONNX direct indisponible ({e}) -> conversion PyTorch")
# (b) Téléchargement des fichiers source puis conversion PyTorch -> ONNX
src = os.path.join(_MODELS, "_ai_src")
os.makedirs(src, exist_ok=True)
hf_hub_download(repo_id=repo, filename="config.json", token=token, local_dir=src)
try:
hf_hub_download(repo_id=repo, filename="preprocessor_config.json", token=token, local_dir=src)
except Exception:
_log("preprocessor_config.json absent — valeurs ImageNet par défaut.")
try:
hf_hub_download(repo_id=repo, filename="pytorch_model.bin", token=token, local_dir=src)
except Exception:
hf_hub_download(repo_id=repo, filename="model.safetensors", token=token, local_dir=src)
_log(f"poids récupérés depuis {repo} -> conversion ONNX en cours…")
_convert_to_onnx(src)
def _convert_to_onnx(src_dir):
import numpy as np
import torch
import torch.nn as nn
from transformers import AutoModelForImageClassification
model, info = AutoModelForImageClassification.from_pretrained(
src_dir, output_loading_info=True, attn_implementation="eager")
miss = [k for k in info.get("missing_keys", [])
if "relative_position_index" not in k and "attn_mask" not in k]
if miss:
raise RuntimeError(f"poids manquants après chargement: {miss[:8]}")
model.eval()
class Wrap(nn.Module):
def __init__(self, m):
super().__init__()
self.m = m
def forward(self, x):
return self.m(pixel_values=x).logits
size = int(getattr(model.config, "image_size", 224) or 224)
dummy = torch.randn(1, 3, size, size)
torch.onnx.export(
Wrap(model).eval(), dummy, AI_ONNX,
input_names=["pixel_values"], output_names=["logits"],
opset_version=17, do_constant_folding=True, dynamo=False)
# Sidecar : labels, taille, normalisation, index « fake »
id2label = {int(k): v for k, v in (model.config.id2label or {0: "fake", 1: "real"}).items()}
def is_fake(name):
n = str(name).lower()
return any(t in n for t in ("artificial", "ai", "fake", "synth", "generated", "diffusion", "midjourney"))
env_idx = os.getenv("SDXL_FAKE_INDEX")
if env_idx not in (None, ""):
fake_index = int(env_idx)
else:
cands = [i for i, n in id2label.items() if is_fake(n)]
fake_index = cands[0] if cands else 0
mean, std, psize = list(_DEF_MEAN), list(_DEF_STD), size
pp = os.path.join(src_dir, "preprocessor_config.json")
if os.path.exists(pp):
try:
d = json.load(open(pp, encoding="utf-8"))
mean = [float(x) for x in d.get("image_mean", mean)]
std = [float(x) for x in d.get("image_std", std)]
s = d.get("size", size)
if isinstance(s, dict):
psize = int(s.get("height") or s.get("shortest_edge") or size)
else:
psize = int(s)
except Exception:
pass
sidecar = {"size": int(psize), "mean": mean, "std": std,
"labels": {str(k): v for k, v in id2label.items()},
"fake_index": int(fake_index)}
with open(AI_SIDECAR, "w", encoding="utf-8") as f:
json.dump(sidecar, f, ensure_ascii=False, indent=2)
_log(f"ai_detector.onnx prêt (labels={id2label}, fake_index={fake_index}, size={psize}).")
# ── 2. Modèle visage (EfficientNet / FaceForensics++) ───────────────────────
def ensure_face_model():
if os.path.exists(FACE_ONNX):
_log("deepfake_efficientnet.onnx déjà présent — rien à télécharger.")
return
explicit = os.getenv("ONNX_MODEL_PATH")
if explicit and os.path.exists(explicit):
_log(f"ONNX_MODEL_PATH fourni ({explicit}) — pas de téléchargement.")
return
repo = os.getenv("ONNX_MODEL_REPO")
fname = os.getenv("ONNX_MODEL_FILE")
if not repo or not fname:
_log("modèle visage non téléchargé : définir ONNX_MODEL_REPO + ONNX_MODEL_FILE "
"(repo HF + fichier .onnx). L'app fonctionne sans (analyse visage désactivée).")
return
from huggingface_hub import hf_hub_download
os.makedirs(_MODELS, exist_ok=True)
f = hf_hub_download(repo_id=repo, filename=fname, token=_hf_token())
shutil.copy(f, FACE_ONNX)
_log(f"deepfake_efficientnet.onnx téléchargé depuis {repo}/{fname}")
# ── Modèle vidéo ResNext50+LSTM ─────────────────────────────────────────────
# Téléchargé ET chargé PARESSEUSEMENT par video_model.VideoDeepfakeAnalyzer, au
# tout premier appel à predict() (première analyse vidéo) — plus au démarrage.
# ── 3. Détecteur de visages YuNet (famille RetinaFace, ~230 Ko) ─────────────
def ensure_yunet():
if os.path.exists(YUNET_ONNX):
_log("face_detector_yunet.onnx déjà présent — rien à télécharger.")
return
import urllib.request
os.makedirs(_MODELS, exist_ok=True)
tmp = YUNET_ONNX + ".part"
urllib.request.urlretrieve(YUNET_URL, tmp)
if os.path.getsize(tmp) < 100_000: # un pointeur LFS/HTML fait < 100 Ko
os.unlink(tmp)
raise RuntimeError("téléchargement YuNet invalide (fichier trop petit)")
os.replace(tmp, YUNET_ONNX)
_log("face_detector_yunet.onnx téléchargé (YuNet).")
# ── Orchestration arrière-plan ──────────────────────────────────────────────
def _run(detector):
for label, fn in (("détecteur de visages", ensure_yunet),
("IA global", ensure_sdxl_model), ("visage", ensure_face_model)):
try:
fn()
except Exception as e:
_log(f"échec modèle {label} : {e}")
# Les analyseurs chargent leur ONNX paresseusement : comme le fichier était
# absent au démarrage, `_tried` est resté None et le prochain `predict()`
# rechargera automatiquement le modèle fraîchement téléchargé. Rien à forcer.
_log("provisionnement terminé.")
def start_background_bootstrap(detector=None):
"""Lance le provisionnement des modèles dans un thread démon (non bloquant)."""
if os.getenv("DISABLE_MODEL_BOOTSTRAP", "0").lower() in ("1", "true", "yes"):
_log("désactivé (DISABLE_MODEL_BOOTSTRAP).")
return
threading.Thread(target=_run, args=(detector,), name="model-bootstrap", daemon=True).start()
_log("démarré (arrière-plan).")