Whyx-PROmpTea / src /moondream_tagger.py
ArtShumov's picture
feat(tagger): swap Qwen-VL for Moondream2 NL captioner
17a1e39
Raw
History Blame
5.61 kB
"""Moondream2 captioner for natural-language scene descriptions.
Moondream2 is a much smaller VLM than Qwen3-VL (~3.7 GB bf16, ~2B params) and
runs well on CPU-only HF Spaces free tier (2 vCPU / 16 GB RAM). The model card
states the public API over the *model class itself* — caption / query / detect /
point — so we call the high-level helpers directly instead of a generic
"generate()" loop.
By default the model is loaded lazily on first call (CPU). The runtime RAM
footprint stays well under 8 GB when stacked alongside WD14 taggers, pose
estimators and the SD pipeline.
Set ``WHYX_DISABLE_MOONDREAM=1`` to keep Moondream off in constrained
deployments.
"""
from __future__ import annotations
import io
import logging
import os
import threading
from typing import Optional
logger = logging.getLogger(__name__)
try:
import torch
from PIL import Image, ImageOps
from transformers import AutoModelForCausalLM, AutoTokenizer
except Exception: # pragma: no cover - keep import errors soft
AutoModelForCausalLM = AutoTokenizer = None
torch = None
Image = ImageOps = None
_MODEL_ID = os.environ.get("WHYX_MOONDREAM_MODEL", "vikhyatk/moondream2")
_MD_INSTANCE: "MoondreamTagger | None" = None
def _md_enabled() -> bool:
return os.environ.get("WHYX_DISABLE_MOONDREAM", "0").strip().lower() not in ("1", "true", "yes", "on")
class MoondreamTagger:
"""One-stop wrapper around `vikhyatk/moondream2` for caption + query.
"""
def __init__(self, model_id: str = _MODEL_ID):
self._model_id = model_id
self._model: Optional[AutoModelForCausalLM] = None
self._tokenizer: Optional[AutoTokenizer] = None
self._loaded = False
self._lock = threading.Lock()
def ensure_loaded(self) -> bool:
if self._loaded:
return True
if not _md_enabled() or AutoModelForCausalLM is None:
return False
with self._lock:
if self._loaded:
return True
try:
# bf16 halves RAM vs fp32 when a GPU is present; CPU gets fp32
# (still well under the 16 GB cap once the model is alive).
dtype = torch.bfloat16 if torch.cuda.is_available() else torch.float32
self._model = AutoModelForCausalLM.from_pretrained(
self._model_id,
dtype=dtype,
# moondream2's anti-hallucination / EOS logic lives in its config
trust_remote_code=True,
# Keep all RAM on CPU; avoid any accelerate/gpu dispatch.
device_map={"": "cpu"},
)
self._tokenizer = AutoTokenizer.from_pretrained(self._model_id)
except Exception as exc: # pragma: no cover
logger.warning("Moondream2 failed to load: %s", exc)
return False
self._loaded = True
return True
@staticmethod
def _to_pil(image) -> "Image.Image":
# Accept PIL/numpy/raw bytes; EXIF orientation is normalised.
if Image is None:
raise RuntimeError("PIL is not available")
if isinstance(image, Image.Image):
pil = image.convert("RGB")
else:
try:
import numpy as np
if isinstance(image, np.ndarray):
pil = Image.fromarray(image)
else:
pil = Image.open(io.BytesIO(image)).convert("RGB")
except Exception as exc:
raise ValueError(f"cannot convert input to PIL image: {exc}") from exc
if ImageOps is not None:
pil = ImageOps.exif_transpose(pil)
return pil
def caption(self, image, length: str = "normal") -> str:
"""Return a natural-language caption for the image.
``length`` may be ``"short"`` (one phrase) or ``"normal"`` (one sentence).
"""
if not self.ensure_loaded():
return ""
pil = self._to_pil(image)
try:
result = self._model.caption(pil, length=length)
except Exception as exc: # pragma: no cover
logger.warning("Moondream2 caption failed: %s", exc)
return ""
# The public API returns a dict with a "caption" key. Helper text is
# stripped to one sentence (and one line) so it can join tags cleanly.
text = result.get("caption", "")
return " ".join(text.split())
def query(self, image, question: str) -> str:
"""Free-form visual Q&A. E.g. "How many people are in the image?" """
if not self.ensure_loaded():
return ""
pil = self._to_pil(image)
try:
result = self._model.query(pil, question)
except Exception as exc: # pragma: no cover
logger.warning("Moondream2 query failed: %s", exc)
return ""
return " ".join(result.get("answer", "").split())
def detect(self, image, thing: str) -> int:
"""Detect instances of `thing` in `image`, return count."""
if not self.ensure_loaded():
return 0
pil = self._to_pil(image)
try:
result = self._model.detect(pil, thing)
except Exception:
return 0
return len(result.get("objects", []))
import io
import threading
def get_moondream_tagger() -> MoondreamTagger:
"""Singleton accessor shared across calls (one model per process)."""
global _MD_INSTANCE
if _MD_INSTANCE is None:
_MD_INSTANCE = MoondreamTagger()
return _MD_INSTANCE