File size: 5,921 Bytes
c881b77 | 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 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 | from . import transforms
from .ref_table import RefTable
from utils import get_ckpt_path
import databuilders
import datasets
import torch
import json
import os
class Loader:
def __init__(self, config, image_processor, split='train', logger=None):
self.logger = logger
self.split, self.phase = split, config.phase
self.data_name = config.dataset.name
# self.data_files = config.dataset.data_files
self.data_files = config.dataset.data_files[self.split][self.phase]
self.categories = config.dataset.categories[self.phase]
self.image_column, self.caption_column, self.bbox_column, self.obbox_column = config.dataset.column_names
self.batch_size = config.training.batch_size
self.num_workers = config.training.num_workers
self.max_inference_size = config.inference.get('max_inference_size', None)
self.novel_sample_dict = None
self.set_dataset()
if self.phase == 'novel':
self.k_shot = config.dataset.novel_settings.k_shot
self.shuffle_seed = config.dataset.novel_settings.get('shuffle_seed', 42)
self.dump_file = os.path.join(get_ckpt_path(config), 'novel_sample_dict.json')
if self.split == 'train':
self.sample_dataset()
elif self.split == 'infer':
self.load_novel_sample_dict()
self.sample_dataset_infer_phase()
self.dataset = datasets.concatenate_datasets(self.dataset.values())
self.ref_table = RefTable(config, filter_dict=self.novel_sample_dict)
self.transform = getattr(transforms, config.dataset.get('transform', 'DefaultTransform'))(config, image_processor, self.split, ref_table=self.ref_table())
self.dataset = self.dataset.with_transform(self.transform)
def set_dataset(self):
builder = getattr(databuilders, self.data_name.lower(), None)
if builder is None:
raise ValueError(f"Unknown dataset: {self.data_name}")
builder = builder(data_files=self.data_files)
builder.download_and_prepare()
self.dataset = builder.as_dataset()
def sample_dataset(self):
''' The data sample logics for few-shot learning. '''
self.novel_sample_dict = {}
for category in self.categories:
self.dataset[category] = self.dataset[category].shuffle(self.shuffle_seed).select(range(min(self.k_shot, len(self.dataset[category]))))
self.novel_sample_dict[category] = list(self.dataset[category]['dataid'])
def sample_dataset_infer_phase(self):
if self.max_inference_size is not None:
# self.dataset['default'] = self.dataset['default'].shuffle(self.shuffle_seed).select(range(self.max_inference_size))
for category in self.categories:
self.dataset[category] = self.dataset[category].shuffle(self.shuffle_seed).select(range(min(self.max_inference_size, len(self.dataset[category]))))
def collate_fn(self, examples):
images = torch.stack([example[self.image_column] for example in examples])
images = images.to(memory_format=torch.contiguous_format).float()
captions = [example[self.caption_column] for example in examples]
bboxes = [example[self.bbox_column] for example in examples]
obboxes = [example[self.obbox_column] for example in examples]
if isinstance(examples[0]["instances"], list):
instances = [example["instances"] for example in examples]
else:
instances = torch.stack([example["instances"] for example in examples])
dataid = [example["dataid"] for example in examples]
# Custom Keys
if 'masks' in examples[0].keys():
masks = torch.stack([example["masks"] for example in examples])
return {self.image_column: images, self.caption_column: captions, self.bbox_column: bboxes, self.obbox_column: obboxes, "instances": instances, "dataid": dataid, "masks": masks}
if 'parallels' in examples[0].keys():
# parallels = torch.cat([example["parallels"] for example in examples])
# images = torch.cat([images, parallels], dim=0)
parallels = [example["parallels"] for example in examples]
return {self.image_column: images, self.caption_column: captions, self.bbox_column: bboxes, self.obbox_column: obboxes, "instances": instances, "dataid": dataid, "parallels": parallels}
return {self.image_column: images, self.caption_column: captions, self.bbox_column: bboxes, self.obbox_column: obboxes, "instances": instances, "dataid": dataid}
def __call__(self):
# for i in range(6):
# _ = self.dataset[i]
dataloader = torch.utils.data.DataLoader(
self.dataset,
shuffle=True,
collate_fn=self.collate_fn,
batch_size=self.batch_size,
num_workers=self.num_workers,
)
return dataloader
def dump_novel_sample_dict(self):
if self.phase == 'novel' and self.logger is not None:
self.logger.info(f'Novel Sample Dict: \n{json.dumps(self.novel_sample_dict, indent=2)}', main_process_only=False)
self.logger.info(f'Novel Ref Table: \n{json.dumps(self.ref_table("novel"), indent=2)}', main_process_only=False)
self.logger.info(f'Dump novel_sample_dict at {self.dump_file}')
with open(self.dump_file, 'w') as f:
json.dump(self.novel_sample_dict, f, indent=2)
def load_novel_sample_dict(self):
if self.phase == 'novel' and self.logger is not None:
self.logger.info(f'Load novel_sample_dict at {self.dump_file}')
with open(self.dump_file, 'r') as f:
self.novel_sample_dict = json.load(f)
self.logger.info(f'Novel Sample Dict: \n{json.dumps(self.novel_sample_dict, indent=2)}', main_process_only=False) |