| """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: |
| 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: |
| import traceback |
|
|
| traceback.print_exc() |
| yield f"Error: {e!s}" |
|
|
|
|