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
|