Deploybot
Deploy from stable branch
082393b
Raw
History Blame Contribute Delete
4.04 kB
"""Chart-to-CSV extraction using Granite Vision.
Converts chart images to tabular CSV data using ibm-granite/granite-vision-4.1-4b.
"""
import threading
from collections.abc import Generator
from PIL import Image
from typing import Any
import torch
from PIL import Image
from transformers import TextIteratorStreamer
from model_loader import load_model, use_api_mode, use_mlx_mode
def extract_csv(image: Image.Image) -> str:
"""Extract CSV data from a chart image using Granite Vision.
Args:
image: PIL Image of a chart or table.
Returns:
CSV-formatted text extracted from the chart.
"""
if use_api_mode():
from infer_api import extract_csv_api
return extract_csv_api(image)
if use_mlx_mode():
from infer_mlx import extract_csv_mlx
return extract_csv_mlx(image)
processor, model = load_model()
if processor is None or model is None:
return "col1,col2,col3\nvalue1,value2,value3\nvalue4,value5,value6"
try:
import torch
image = image.convert("RGB")
conversation = [{"role": "user", "content": [
{"type": "image"},
{"type": "text", "text": "<chart2csv>"},
]}]
text = processor.apply_chat_template(conversation, tokenize=False, add_generation_prompt=True)
inputs = processor(text=text, images=image, return_tensors="pt").to(model.device)
max_new_tokens = 4096
with torch.inference_mode():
outputs = model.generate(**inputs, max_new_tokens=max_new_tokens, use_cache=True)
gen = outputs[0, inputs["input_ids"].shape[1]:]
result = processor.decode(gen, skip_special_tokens=True)
if len(gen) >= max_new_tokens:
result += "\n\n[Max token limit reached — response may be truncated]"
return result
except Exception as e: # noqa: BLE001
import traceback
traceback.print_exc()
return f"Error: {e!s}"
def _run_generate(model: Any, generation_kwargs: dict[str, Any]) -> None:
"""Run model.generate in a thread (used by the streaming variant)."""
with torch.inference_mode():
model.generate(**generation_kwargs)
def extract_csv_stream(image: Image.Image) -> Generator[str, None, None]:
"""Stream CSV extraction token-by-token from a chart image.
Same interface as extract_csv() but yields tokens incrementally.
Args:
image: PIL Image of a chart or table.
Yields:
Token strings as they are generated by the model.
"""
if use_api_mode():
from infer_api import extract_csv_stream_api
yield from extract_csv_stream_api(image)
return
if use_mlx_mode():
from infer_mlx import extract_csv_stream_mlx
yield from extract_csv_stream_mlx(image)
return
processor, model = load_model()
if processor is None or model is None:
yield "col1,col2,col3\nvalue1,value2,value3\nvalue4,value5,value6"
return
try:
image = image.convert("RGB")
conversation = [{"role": "user", "content": [
{"type": "image"},
{"type": "text", "text": "<chart2csv>"},
]}]
text = processor.apply_chat_template(conversation, tokenize=False, add_generation_prompt=True)
inputs = processor(text=text, images=image, return_tensors="pt").to(model.device)
streamer = TextIteratorStreamer(
processor.tokenizer, skip_prompt=True, skip_special_tokens=True,
)
generation_kwargs = {
**inputs,
"max_new_tokens": 4096,
"use_cache": True,
"streamer": streamer,
}
thread = threading.Thread(target=_run_generate, args=(model, generation_kwargs))
thread.start()
for token_text in streamer:
if token_text:
yield token_text
thread.join()
except Exception as e: # noqa: BLE001
import traceback
traceback.print_exc()
yield f"Error: {e!s}"