Studio: caption RAG figures via loaded chat VLM only; drop separate captioning model

This commit is contained in:
Roland Tannous 2026-05-26 23:53:23 +04:00
commit c6a1935dd2
2 changed files with 136 additions and 116 deletions

View file

@ -1,146 +1,109 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Figure captioning via a small generative VLM (PaddleOCR-VL, 4-bit).
"""Figure captioning via the user's currently-loaded chat VLM.
Defensive by design: any load or per-image failure returns an empty
string so the ingestion pipeline can fall back to its prior page-text
caption behaviour rather than crashing.
No separate vision model is loaded. If the chat model is vision-capable
(detected at ingestion-enqueue time and passed in as ``vlm_url`` /
``vlm_model``), we call its OpenAI-compatible ``/v1/chat/completions``
with the figure as a base64 ``image_url``. If no vision-capable chat
model is loaded, captioning is skipped entirely the ingestion falls
back to the parser's page-text ``nearest_caption``.
Lifecycle: lives inside the ingestion subprocess (`_subprocess_worker`
in `ingestion.py`). The model loads on first `caption_images` call and
is released when the subprocess exits no impact on the parent chat
model.
Defensive: any per-image request failure returns an empty string so
ingestion stays resilient.
"""
from __future__ import annotations
import base64
import logging
import threading
from typing import Any
from io import BytesIO
from typing import Optional
logger = logging.getLogger(__name__)
# Pre-quantized Unsloth bnb-4bit repo (see unsloth/models/mapper.py).
# Using this name directly skips the FLOAT_TO_INT_MAPPER redirect so the
# initial download fetches the 4-bit weights (~1.5 GB) instead of bf16
# (~5 GB followed by on-the-fly quantization).
_CAPTION_MODEL_NAME = "unsloth/Qwen3-VL-2B-Instruct-unsloth-bnb-4bit"
_PROMPT = (
"Describe this figure in <=60 words. Focus on factual content "
"(axes, labels, captions, visible text, main objects). "
"Do not speculate beyond what is visible."
)
_MAX_NEW_TOKENS = 120
_lock = threading.Lock()
_model: Any | None = None
_processor: Any | None = None
_load_failed: bool = False
# Downscale large images so the base64 payload stays manageable; the chat
# model's prefill cost scales with image-tile count, not pixel count, but
# very large inputs still bloat the JSON body. 1600 px on the long side
# matches PR #5351's chat-composer extractor.
_MAX_IMAGE_SIZE = 1600
_REQUEST_TIMEOUT_SECONDS = 120.0
def _load() -> tuple[Any, Any] | None:
"""Lazy-load PaddleOCR-VL. Returns (model, processor) or None on failure.
def _image_to_data_url(blob: bytes) -> str:
from PIL import Image
Sentinel-cached: once a load failure occurs in this subprocess we
don't keep retrying for every batch of images.
"""
global _model, _processor, _load_failed
if _load_failed:
return None
with _lock:
if _model is not None and _processor is not None:
return _model, _processor
try:
from unsloth import FastVisionModel
logger.info(
"Loading RAG captioner: %s (4-bit via Unsloth)",
_CAPTION_MODEL_NAME,
)
# FastVisionModel handles 4-bit quantization, dtype selection,
# and inference-mode wiring. Returns (model, processor) — for
# vision models the "tokenizer" slot carries the processor.
model, processor = FastVisionModel.from_pretrained(
model_name = _CAPTION_MODEL_NAME,
load_in_4bit = True,
trust_remote_code = True,
)
FastVisionModel.for_inference(model)
_model = model
_processor = processor
return _model, _processor
except Exception as exc:
logger.warning(
"RAG captioner %s failed to load (%s). Falling back to "
"page-text captions for this ingestion.",
_CAPTION_MODEL_NAME,
exc,
)
_load_failed = True
return None
img = Image.open(BytesIO(blob)).convert("RGB")
if max(img.size) > _MAX_IMAGE_SIZE:
img.thumbnail((_MAX_IMAGE_SIZE, _MAX_IMAGE_SIZE))
buf = BytesIO()
img.save(buf, format = "JPEG", quality = 88)
encoded = base64.b64encode(buf.getvalue()).decode("ascii")
return f"data:image/jpeg;base64,{encoded}"
def caption_images(image_bytes_list: list[bytes]) -> list[str]:
"""Generate one short caption per image. Same-length output.
def caption_images(
image_bytes_list: list[bytes],
*,
vlm_url: Optional[str] = None,
vlm_model: Optional[str] = None,
) -> list[str]:
"""Caption each image via the loaded chat VLM.
Returns ``""`` for any image whose captioning failed (or all images
if the model couldn't load). Never raises — ingestion must stay
resilient to VLM unavailability.
``vlm_url`` / ``vlm_model`` come from the parent's chat-backend probe.
When either is missing (no model loaded, or loaded model is text-only),
returns an empty string per image so the caller falls back to its
parser-provided caption. Never raises.
"""
if not image_bytes_list:
return []
loaded = _load()
if loaded is None:
if not vlm_url or not vlm_model:
return ["" for _ in image_bytes_list]
model, processor = loaded
from io import BytesIO
import torch
from PIL import Image
import httpx
endpoint = f"{vlm_url.rstrip('/')}/v1/chat/completions"
out: list[str] = []
for blob in image_bytes_list:
try:
image = Image.open(BytesIO(blob)).convert("RGB")
messages = [
{
"role": "user",
"content": [
{"type": "image", "image": image},
{"type": "text", "text": _PROMPT},
with httpx.Client(timeout = _REQUEST_TIMEOUT_SECONDS) as client:
for blob in image_bytes_list:
try:
data_url = _image_to_data_url(blob)
payload = {
"model": vlm_model,
"messages": [
{
"role": "user",
"content": [
{"type": "text", "text": _PROMPT},
{
"type": "image_url",
"image_url": {"url": data_url},
},
],
}
],
"max_tokens": _MAX_NEW_TOKENS,
"temperature": 0.0,
}
]
# Qwen3-VL's processor pipes image + text through the chat
# template in one call when tokenize=True / return_dict=True.
inputs = processor.apply_chat_template(
messages,
tokenize = True,
add_generation_prompt = True,
return_dict = True,
return_tensors = "pt",
)
inputs = {
k: (v.to(model.device) if hasattr(v, "to") else v)
for k, v in inputs.items()
}
input_ids_len = int(inputs["input_ids"].shape[1])
with torch.no_grad():
output_ids = model.generate(
**inputs,
max_new_tokens = _MAX_NEW_TOKENS,
do_sample = False,
response = client.post(endpoint, json = payload)
response.raise_for_status()
data = response.json()
content = (
data.get("choices", [{}])[0].get("message", {}).get("content", "")
)
# Decode only the newly-generated tokens, not the prompt echo.
new_tokens = output_ids[:, input_ids_len:]
caption = processor.batch_decode(
new_tokens,
skip_special_tokens = True,
)[0]
out.append(caption.strip())
except Exception as exc: # noqa: BLE001
logger.warning("caption_images: skipping one image: %s", exc)
out.append("")
out.append(content.strip() if isinstance(content, str) else "")
except Exception as exc: # noqa: BLE001
logger.warning(
"caption_images: per-image request to %s failed: %s",
endpoint,
exc,
)
out.append("")
return out

View file

@ -60,6 +60,8 @@ def _subprocess_worker(
chunking_strategy: str = "standard",
mode: str = "text",
document_id: str = "",
vlm_url: str | None = None,
vlm_model: str | None = None,
) -> None:
try:
from core.rag.chunking import chunk_pages
@ -117,6 +119,8 @@ def _subprocess_worker(
model_name = model_name,
out_queue = out_queue,
first_index = text_count,
vlm_url = vlm_url,
vlm_model = vlm_model,
)
out_queue.put({"type": "complete", "num_chunks": text_count + image_count})
except Exception as exc: # noqa: BLE001
@ -190,6 +194,8 @@ def _stream_image_chunks(
model_name: str,
out_queue,
first_index: int,
vlm_url: str | None = None,
vlm_model: str | None = None,
) -> int:
"""Persist images, emit image+caption chunks; pairs share pair_group."""
from core.rag.captioner import caption_images
@ -223,10 +229,15 @@ def _stream_image_chunks(
if not paths:
return 0
# VLM-generated captions; falls back to the parser's nearest_caption
# (page-text blob) when the VLM is unavailable or fails for an image.
# VLM-generated captions via the loaded chat VLM (when one is loaded
# and vision-capable). Falls back to the parser's nearest_caption
# (page-text blob) when no VLM is available or a request fails.
out_queue.put({"type": "progress", "stage": "caption_images", "progress": 0.87})
vlm_captions = caption_images(bytes_for_encoding)
vlm_captions = caption_images(
bytes_for_encoding,
vlm_url = vlm_url,
vlm_model = vlm_model,
)
captions: list[str] = [
(
vlm_captions[i].strip()
@ -638,6 +649,34 @@ def _pump(
state.push_event({"type": "error", "error": final_error})
def _probe_loaded_vlm() -> tuple[str | None, str | None]:
"""Best-effort: return (base_url, model_name) when a vision-capable
chat model is currently loaded via llama-server; (None, None) otherwise.
Used to caption figures with the user's chat VLM instead of loading
a dedicated captioning model no extra VRAM, no extra download.
Only the llama-server backend is supported today; transformers /
unsloth in-process VLMs would need a different bridge.
"""
try:
from core.inference.llama_cpp import get_llama_cpp_backend
except Exception:
return None, None
try:
backend = get_llama_cpp_backend()
except Exception:
return None, None
if not getattr(backend, "is_loaded", False):
return None, None
if not getattr(backend, "is_vision", False):
return None, None
base_url = getattr(backend, "base_url", None)
model_id = getattr(backend, "model_identifier", None)
if not base_url or not model_id:
return None, None
return base_url, model_id
def enqueue_ingestion(
document_id: str,
stored_path: Path,
@ -657,6 +696,22 @@ def enqueue_ingestion(
or resolve_embedder(mode, chunking_strategy)
or RAG_EMBEDDING_MODEL
)
# Probe before forking the subprocess so the loaded-model info is
# captured in the parent's process state, then passed to the child.
vlm_url, vlm_model = (None, None)
if mode == "multimodal":
vlm_url, vlm_model = _probe_loaded_vlm()
if vlm_url:
logger.info(
"RAG ingest: will caption figures via loaded chat VLM %s at %s",
vlm_model,
vlm_url,
)
else:
logger.info(
"RAG ingest: no vision-capable chat model loaded; "
"skipping figure captioning (fallback to page-text)."
)
job_id = str(uuid4())
with get_connection() as conn:
conn.execute(
@ -686,6 +741,8 @@ def enqueue_ingestion(
chunking_strategy,
mode,
document_id,
vlm_url,
vlm_model,
),
daemon = True,
)