File size: 4,164 Bytes
c6bc767
 
 
 
 
 
9a1a7d2
92c7321
c6bc767
 
 
9a1a7d2
c6bc767
fe5285c
9a1a7d2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c6bc767
 
 
 
 
 
 
 
9a1a7d2
c6bc767
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9a1a7d2
c6bc767
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
83
84
85
86
87
88
89
90
91
92
93
import io
import json
import torch
from torch.utils.data import Dataset
from PIL import Image

from utils import Logger
from dataset.common import VLM, pre_processing_chat, post_processing_chat


class VLMDataset(Dataset):
    def __init__(self, parquet_path, tokenizer, preprocess=None, max_length=512, image_special_token='<|image_pad|>', image_token_len=64, max_samples=None):
        super().__init__()
        import pyarrow.parquet as pq
        pf = pq.ParquetFile(parquet_path)
        total = pf.metadata.num_rows
        if max_samples is not None and max_samples < total:
            total = max_samples
        cols = pf.metadata.schema.names
        rows = {c: [] for c in cols}
        loaded = 0
        for batch in pf.iter_batches(batch_size=32768, columns=cols):
            batch = batch.slice(0, total - loaded)
            for c in cols:
                rows[c].extend(batch.column(c).to_pylist())
            loaded += batch.num_rows
            if loaded >= total:
                break
        self.data = rows
        del pf, rows
        Logger(f'Loaded {loaded} samples from {parquet_path}')
        self.tokenizer = tokenizer
        self.max_length = max_length
        self.preprocess = preprocess
        self.image_special_token = image_special_token * image_token_len
        self.bos_id = tokenizer(f'{tokenizer.bos_token}assistant\n', add_special_tokens=False).input_ids
        self.eos_id = tokenizer(f'{tokenizer.eos_token}\n', add_special_tokens=False).input_ids

    def __len__(self):
        return len(self.data['conversations']) if 'conversations' in self.data else 0

    def create_chat_prompt(self, conversations):
        messages = []
        for turn in conversations:
            content = turn['content'].replace('<image>', self.image_special_token) if turn.get('role') != 'system' else turn['content']
            messages.append({"role": turn['role'], "content": content})
        tools = conversations[0]["functions"] if (conversations and conversations[0]["role"] == "system" and conversations[0].get("functions")) else None
        return self.tokenizer.apply_chat_template(
            messages,
            tokenize=False,
            add_generation_prompt=False,
            tools=tools
        )

    def generate_labels(self, input_ids):
        labels = [-100] * len(input_ids)
        i = 0
        while i < len(input_ids):
            if input_ids[i:i + len(self.bos_id)] == self.bos_id:
                start = i + len(self.bos_id)
                end = start
                while end < len(input_ids):
                    if input_ids[end:end + len(self.eos_id)] == self.eos_id:
                        break
                    end += 1
                for j in range(start, min(end + len(self.eos_id), self.max_length)):
                    labels[j] = input_ids[j]
                i = end + len(self.eos_id) if end < len(input_ids) else len(input_ids)
            else:
                i += 1
        return labels

    def __getitem__(self, index: int):
        row = {k: v[index] for k, v in self.data.items()}
        conversations = json.loads(row['conversations']) if isinstance(row['conversations'], str) else row['conversations']
        image_bytes = row['image_bytes']
        if not isinstance(image_bytes, list): image_bytes = [image_bytes]

        conversations = pre_processing_chat(conversations)
        prompt = self.create_chat_prompt(conversations)
        prompt = post_processing_chat(prompt)
        input_ids = self.tokenizer(prompt).input_ids[:self.max_length]
        input_ids += [self.tokenizer.pad_token_id] * (self.max_length - len(input_ids))
        labels = self.generate_labels(input_ids)

        image_inputs_list = [VLM.image2tensor(Image.open(io.BytesIO(img)), self.preprocess) for img in image_bytes]
        if hasattr(image_inputs_list[0], 'keys'):
            image_data = {k: torch.cat([inp[k] for inp in image_inputs_list], dim=0) for k in image_inputs_list[0].keys()}
        else:
            image_data = torch.stack(image_inputs_list)

        return torch.tensor(input_ids, dtype=torch.long), torch.tensor(labels, dtype=torch.long), image_data