Studio: caption RAG figures via loaded chat VLM only; drop separate captioning model
This commit is contained in:
parent
f522545b65
commit
c6a1935dd2
2 changed files with 136 additions and 116 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue