File size: 3,007 Bytes
be90b31
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
"""
Dataset class for Sci-Image scientific figure to markdown table task.
"""

import os
from typing import Any, Dict, List, Optional
from PIL import Image
import torch
from torch.utils.data import Dataset

from ..utils.io import read_jsonl


class SciImageTableDataset(Dataset):
    """PyTorch Dataset for scientific figure panels and markdown table targets."""

    def __init__(
        self,
        data_path: str,
        image_dir: Optional[str] = None,
        system_prompt: str = "Extract the plotted data from this figure into a clean Markdown table.",
        max_image_resolution: Optional[int] = None,
        **kwargs: Any,
    ):
        self.data_path = data_path
        self.image_dir = image_dir
        self.system_prompt = system_prompt
        self.max_image_resolution = max_image_resolution

        if not os.path.exists(data_path):
            raise FileNotFoundError(f"Dataset file not found: {data_path}")

        self.records: List[Dict[str, Any]] = read_jsonl(data_path)

    def __len__(self) -> int:
        return len(self.records)

    def _resolve_image_path(self, img_path: str) -> str:
        if os.path.isabs(img_path) and os.path.exists(img_path):
            return img_path
        if self.image_dir:
            candidate = os.path.join(self.image_dir, img_path)
            if os.path.exists(candidate):
                return candidate
        # Check relative to dataset file directory
        parent_dir = os.path.dirname(os.path.abspath(self.data_path))
        candidate = os.path.join(parent_dir, img_path)
        if os.path.exists(candidate):
            return candidate
        # Check raw sibling directory or basename fallbacks
        raw_candidate = os.path.join(os.path.dirname(parent_dir), "raw", img_path)
        if os.path.exists(raw_candidate):
            return raw_candidate
        basename_candidate = os.path.join(parent_dir, "images", os.path.basename(img_path))
        if os.path.exists(basename_candidate):
            return basename_candidate
        raw_basename_candidate = os.path.join(os.path.dirname(parent_dir), "raw", "images", os.path.basename(img_path))
        if os.path.exists(raw_basename_candidate):
            return raw_basename_candidate
        return img_path

    def __getitem__(self, idx: int) -> Dict[str, Any]:
        item = self.records[idx]

        image_path = self._resolve_image_path(item["image"])
        image = Image.open(image_path).convert("RGB")
        if self.max_image_resolution:
            image.thumbnail((self.max_image_resolution, self.max_image_resolution))

        markdown_table = item.get("table", item.get("markdown", item.get("target", "")))
        sample_id = item.get("id", str(idx))
        metadata = item.get("metadata", {})

        return {
            "id": sample_id,
            "image": image,
            "image_path": image_path,
            "table": markdown_table,
            "system_prompt": self.system_prompt,
            "metadata": metadata,
        }