from torch.utils.data import Dataset, DataLoader from functools import partial # A synthetic Dataset with random images and captions class FakeDataset(Dataset): def __init__(self): self.length_dist = { 'short': {'range': (5, 32), 'weight': 0.90}, 'medium': {'range': (33, 64), 'weight': 0.09}, 'long': {'range': (65, 256), 'weight': 0.01} } super().__init__() def __len__(self): return 1000000 def __getitem__(self, index): length_bin = random.choices( list(self.length_dist.keys()), weights=[d['weight'] for d in self.length_dist.values()], k=1 )[0] range_start, range_end = self.length_dist[length_bin]['range'] image = torch.randn(3, 224, 224) length = random.randint(range_start, range_end - 1) labels = torch.cat([torch.randint(1, NUM_TOKENS, (length,)), torch.tensor([END_ID])], dim=0) input_ids = torch.cat([torch.tensor([START_ID]), labels[:-1]], dim=0) return { 'image': image, 'input_ids': input_ids, 'labels': labels } def pad_sequence(sequence, length, pad_val): return torch.nn.functional.pad( sequence, (0, length - sequence.shape[0]), value=pad_val ) def collate_with_padding(batch, pad_to_longest=False, align=None): padded_inputs = [] padded_labels = [] if pad_to_longest: pad_len = max([b['input_ids'].shape[0] for b in batch]) if align: pad_len = ((pad_len + align - 1) // align) * align else: pad_len = MAX_SEQ_LEN for b in batch: input_ids = b['input_ids'] labels = b['labels'] padded_inputs.append(pad_sequence(input_ids, pad_len, PAD_ID)) padded_labels.append(pad_sequence(labels, pad_len, -100)) padded_inputs = torch.stack(padded_inputs, dim=0) padded_labels = torch.stack(padded_labels, dim=0) images = torch.stack([b['image'] for b in batch], dim=0) return { 'pixel_values': images, 'decoder_input_ids': padded_inputs, 'labels': padded_labels, 'decoder_attention_mask': (padded_inputs != PAD_ID) } def get_dataloader(pad_to_longest=False, align=None): return DataLoader( dataset=FakeDataset(), batch_size=BATCH_SIZE, num_workers=NUM_WORKERS, collate_fn=partial( collate_with_padding, pad_to_longest=pad_to_longest, align=align ) )