| import ast |
| import io |
| import json |
| import logging |
| import math |
| import os |
| import random |
| import sys |
| import braceexpand |
| from dataclasses import dataclass |
| from multiprocessing import Value |
|
|
| import numpy.lib.format |
| import numpy as np |
| import pandas as pd |
| import torch |
| import torchvision.datasets as datasets |
| import webdataset as wds |
| from PIL import Image |
| from torch.utils.data import Dataset, DataLoader, SubsetRandomSampler, IterableDataset, get_worker_info |
| from torch.utils.data.distributed import DistributedSampler |
| from torchvision.transforms import ToTensor |
| from tqdm import tqdm |
| from webdataset.filters import _shuffle |
| from webdataset.tariterators import base_plus_ext, url_opener, tar_file_expander, valid_sample |
|
|
| from open_clip import get_tokenizer |
| from open_clip.factory import HF_HUB_PREFIX |
| from training.params import parse_args |
| from data.process_text import load_and_transform_text |
| from data.process_video import get_video_transform |
| from data.process_audio import get_audio_transform |
| from data.process_depth import get_depth_transform |
| from data.process_thermal import get_thermal_transform |
| import pdb |
| try: |
| import horovod.torch as hvd |
| except ImportError: |
| hvd = None |
|
|
|
|
|
|
| class SharedEpoch: |
| def __init__(self, epoch: int = 0): |
| self.shared_epoch = Value('i', epoch) |
|
|
| def set_value(self, epoch): |
| self.shared_epoch.value = epoch |
|
|
| def get_value(self): |
| return self.shared_epoch.value |
|
|
|
|
| @dataclass |
| class DataInfo: |
| dataloader: DataLoader |
| sampler: DistributedSampler = None |
| shared_epoch: SharedEpoch = None |
|
|
| def set_epoch(self, epoch): |
| if self.shared_epoch is not None: |
| self.shared_epoch.set_value(epoch) |
| if self.sampler is not None and isinstance(self.sampler, DistributedSampler): |
| self.sampler.set_epoch(epoch) |
|
|
|
|
| def expand_urls(urls, weights=None): |
| if weights is None: |
| expanded_urls = wds.shardlists.expand_urls(urls) |
| return expanded_urls, None |
| if isinstance(urls, str): |
| urllist = urls.split("::") |
| weights = weights.split('::') |
| assert len(weights) == len(urllist), \ |
| f"Expected the number of data components ({len(urllist)}) and weights({len(weights)}) to match." |
| weights = [float(weight) for weight in weights] |
| all_urls, all_weights = [], [] |
| for url, weight in zip(urllist, weights): |
| expanded_url = list(braceexpand.braceexpand(url)) |
| expanded_weights = [weight for _ in expanded_url] |
| all_urls.extend(expanded_url) |
| all_weights.extend(expanded_weights) |
| return all_urls, all_weights |
| else: |
| all_urls = list(urls) |
| return all_urls, weights |
|
|
|
|
| def get_dataset_size(shards): |
| shards_list, _ = expand_urls(shards) |
| dir_path = os.path.dirname(shards_list[0]) |
| sizes_filename = os.path.join(dir_path, 'sizes.json') |
| len_filename = os.path.join(dir_path, '__len__') |
| if os.path.exists(sizes_filename): |
| sizes = json.load(open(sizes_filename, 'r')) |
| total_size = sum([int(sizes[os.path.basename(shard)]) for shard in shards_list]) |
| elif os.path.exists(len_filename): |
| |
| total_size = ast.literal_eval(open(len_filename, 'r').read()) |
| else: |
| total_size = None |
| |
| |
| |
| |
| |
| num_shards = len(shards_list) |
| return total_size, num_shards |
|
|
|
|
|
|
| def count_samples(dataloader): |
| os.environ["WDS_EPOCH"] = "0" |
| n_elements, n_batches = 0, 0 |
| for images, texts in dataloader: |
| n_batches += 1 |
| n_elements += len(images) |
| assert len(images) == len(texts) |
| return n_elements, n_batches |
|
|
|
|
| def filter_no_caption_or_no_image(sample): |
| has_caption = ('raw.txt' in sample and 'mplug.txt' in sample and 'polish_mplug.txt' in sample and 'ofa3.txt' in sample) |
| has_image = ('frm7.jpg' in sample and 'tml0.jpg' in sample and 'dep0.npy' in sample) |
| return has_caption and has_image |
|
|
|
|
| def log_and_continue(exn): |
| """Call in an exception handler to ignore any exception, issue a warning, and continue.""" |
| logging.warning(f'Handling webdataset error ({repr(exn)}). Ignoring.') |
| return True |
|
|
|
|
| def group_by_keys_nothrow(data, keys=base_plus_ext, lcase=True, suffixes=None, handler=None): |
| """Return function over iterator that groups key, value pairs into samples. |
| |
| :param keys: function that splits the key into key and extension (base_plus_ext) |
| :param lcase: convert suffixes to lower case (Default value = True) |
| """ |
| current_sample = None |
| for filesample in data: |
| assert isinstance(filesample, dict) |
| fname, value = filesample["fname"], filesample["data"] |
| prefix, suffix = keys(fname) |
| if prefix is None: |
| continue |
| if lcase: |
| suffix = suffix.lower() |
| |
| |
| |
| if current_sample is None or prefix != current_sample["__key__"] or suffix in current_sample: |
| if valid_sample(current_sample): |
| yield current_sample |
| current_sample = dict(__key__=prefix, __url__=filesample["__url__"]) |
| if suffixes is None or suffix in suffixes: |
| current_sample[suffix] = value |
| if valid_sample(current_sample): |
| yield current_sample |
|
|
|
|
| def tarfile_to_samples_nothrow(src, handler=log_and_continue): |
| |
| streams = url_opener(src, handler=handler) |
| files = tar_file_expander(streams, handler=handler) |
| samples = group_by_keys_nothrow(files, handler=handler) |
| return samples |
|
|
|
|
| def pytorch_worker_seed(increment=0): |
| """get dataloader worker seed from pytorch""" |
| worker_info = get_worker_info() |
| if worker_info is not None: |
| |
| seed = worker_info.seed |
| if increment: |
| |
| seed += increment * max(1, worker_info.num_workers) |
| return seed |
| |
| return wds.utils.pytorch_worker_seed() |
|
|
|
|
| _SHARD_SHUFFLE_SIZE = 200 |
| _SHARD_SHUFFLE_INITIAL = 50 |
| _SAMPLE_SHUFFLE_SIZE = 500 |
| _SAMPLE_SHUFFLE_INITIAL = 100 |
|
|
|
|
| class detshuffle2(wds.PipelineStage): |
| def __init__( |
| self, |
| bufsize=1000, |
| initial=100, |
| seed=0, |
| epoch=-1, |
| ): |
| self.bufsize = bufsize |
| self.initial = initial |
| self.seed = seed |
| self.epoch = epoch |
|
|
| def run(self, src): |
| if isinstance(self.epoch, SharedEpoch): |
| epoch = self.epoch.get_value() |
| else: |
| |
| |
| self.epoch += 1 |
| epoch = self.epoch |
| rng = random.Random() |
| if self.seed < 0: |
| |
| seed = pytorch_worker_seed(epoch) |
| else: |
| |
| seed = self.seed + epoch |
| rng.seed(seed) |
| return _shuffle(src, self.bufsize, self.initial, rng) |
|
|
|
|
| class ResampledShards2(IterableDataset): |
| """An iterable dataset yielding a list of urls.""" |
|
|
| def __init__( |
| self, |
| urls, |
| weights=None, |
| nshards=sys.maxsize, |
| worker_seed=None, |
| deterministic=False, |
| epoch=-1, |
| ): |
| """Sample shards from the shard list with replacement. |
| |
| :param urls: a list of URLs as a Python list or brace notation string |
| """ |
| super().__init__() |
| urls, weights = expand_urls(urls, weights) |
| self.urls = urls |
| self.weights = weights |
| if self.weights is not None: |
| assert len(self.urls) == len(self.weights), \ |
| f"Number of urls {len(self.urls)} and weights {len(self.weights)} should match." |
| assert isinstance(self.urls[0], str) |
| self.nshards = nshards |
| self.rng = random.Random() |
| self.worker_seed = worker_seed |
| self.deterministic = deterministic |
| self.epoch = epoch |
|
|
| def __iter__(self): |
| """Return an iterator over the shards.""" |
| if isinstance(self.epoch, SharedEpoch): |
| epoch = self.epoch.get_value() |
| else: |
| |
| |
| self.epoch += 1 |
| epoch = self.epoch |
| if self.deterministic: |
| |
| if self.worker_seed is None: |
| |
| seed = pytorch_worker_seed(epoch) |
| else: |
| seed = self.worker_seed() + epoch |
| self.rng.seed(seed) |
| for _ in range(self.nshards): |
| if self.weights is None: |
| yield dict(url=self.rng.choice(self.urls)) |
| else: |
| yield dict(url=self.rng.choices(self.urls, weights=self.weights, k=1)[0]) |
|
|
|
|
| class Decode: |
| def __init__(self, args=None): |
| self.num_frames = args.num_frames |
| self.text_type = args.text_type |
| self.chatgpt = self.text_type == 'polish_mplug' |
| self.title = self.text_type == 'raw' |
| self.clip_type = args.clip_type |
| self.tokenizer = get_tokenizer(HF_HUB_PREFIX + args.model, cache_dir=args.cache_dir) |
| self.video_transform = get_video_transform(args) |
| self.audio_transform = get_audio_transform(args) |
| self.depth_transform = get_depth_transform(args) |
| self.thermal_transform = get_thermal_transform(args) |
|
|
|
|
| def __call__(self, sample): |
| input_ids, attention_mask = self.get_text(sample[f"{self.text_type}.txt"], chatgpt=self.chatgpt, title=self.title) |
| if self.clip_type == 'vl': |
| matched_modality = self.get_video([sample[f"frm{i}.jpg"] for i in range(self.num_frames)]) |
| elif self.clip_type == 'al': |
| matched_modality = self.get_audio() |
| elif self.clip_type == 'dl': |
| matched_modality = self.get_depth(sample[f"dep0.npy"]) |
| elif self.clip_type == 'tl': |
| matched_modality = self.get_thermal(sample[f"tml0.jpg"]) |
| |
| else: |
| raise ValueError |
| return matched_modality, input_ids, attention_mask |
|
|
|
|
| def get_video(self, frames): |
| video_data = [] |
| for frame in frames: |
| with io.BytesIO(frame) as stream: |
| img = Image.open(stream) |
| img.load() |
| assert min(img.size) == 256 |
| result = ToTensor()(img) |
| video_data.append(result) |
| video_data = torch.stack(video_data, dim=1) |
| |
| |
| video_outputs = self.video_transform(video_data) |
| return video_outputs |
|
|
|
|
| def get_text(self, text, chatgpt=True, title=False): |
| text = text.decode("utf-8") |
| if chatgpt: |
| assert text.startswith('In the video, ') |
| text = text[14:] |
| tokens = load_and_transform_text(text, self.tokenizer, title=title) |
| return tokens['input_ids'], tokens['attention_mask'] |
|
|
| def get_audio(self): |
| raise NotImplementedError |
|
|
| def get_depth(self, depth): |
| stream = io.BytesIO(depth) |
| img = numpy.lib.format.read_array(stream) |
| depth = self.depth_transform(img) |
| return depth |
|
|
| def get_thermal(self, thermal): |
| with io.BytesIO(thermal) as stream: |
| img = Image.open(stream) |
| img.load() |
| thermal = self.thermal_transform(img) |
| return thermal |
|
|
| def get_wds_dataset(args, is_train, epoch=0, floor=False): |
| input_shards = args.train_data if is_train else args.val_data |
| assert input_shards is not None |
| resampled = getattr(args, 'dataset_resampled', False) and is_train |
|
|
| num_shards = None |
| if is_train: |
| if args.train_num_samples is not None: |
| num_samples = args.train_num_samples |
| else: |
| num_samples, num_shards = get_dataset_size(input_shards) |
| if not num_samples: |
| raise RuntimeError( |
| 'Currently, the number of dataset samples must be specified for the training dataset. ' |
| 'Please specify it via `--train-num-samples` if no dataset length info is present.') |
| else: |
| |
| num_samples = args.val_num_samples or 0 |
|
|
| shared_epoch = SharedEpoch(epoch=epoch) |
|
|
| if resampled: |
| pipeline = [ResampledShards2( |
| input_shards, |
| weights=args.train_data_upsampling_factors, |
| deterministic=True, |
| epoch=shared_epoch, |
| )] |
| else: |
| assert args.train_data_upsampling_factors is None, \ |
| "--train_data_upsampling_factors is only supported when sampling with replacement (with --dataset-resampled)." |
| pipeline = [wds.SimpleShardList(input_shards)] |
|
|
| |
| if is_train: |
| if not resampled: |
| pipeline.extend([ |
| detshuffle2( |
| bufsize=_SHARD_SHUFFLE_SIZE, |
| initial=_SHARD_SHUFFLE_INITIAL, |
| seed=args.seed, |
| epoch=shared_epoch, |
| ), |
| wds.split_by_node, |
| wds.split_by_worker, |
| ]) |
| pipeline.extend([ |
| |
| tarfile_to_samples_nothrow, |
| wds.shuffle( |
| bufsize=_SAMPLE_SHUFFLE_SIZE, |
| initial=_SAMPLE_SHUFFLE_INITIAL, |
| ), |
| ]) |
| else: |
| pipeline.extend([ |
| wds.split_by_worker, |
| |
| wds.tarfile_to_samples(handler=log_and_continue), |
| ]) |
| pipeline.extend([ |
| wds.select(filter_no_caption_or_no_image), |
| |
| |
| |
| |
| wds.map(Decode(args), handler=log_and_continue), |
| wds.batched(args.batch_size, partial=not is_train) |
| ]) |
|
|
| dataset = wds.DataPipeline(*pipeline) |
|
|
| if is_train: |
| if not resampled: |
| num_shards = num_shards or len(expand_urls(input_shards)[0]) |
| assert num_shards >= args.workers * args.world_size, 'number of shards must be >= total workers' |
| |
| round_fn = math.floor if floor else math.ceil |
| global_batch_size = args.batch_size * args.world_size |
| num_batches = round_fn(num_samples / global_batch_size) |
| num_workers = max(1, args.workers) |
| num_worker_batches = round_fn(num_batches / num_workers) |
| num_batches = num_worker_batches * num_workers |
| num_samples = num_batches * global_batch_size |
| dataset = dataset.with_epoch(num_worker_batches) |
| else: |
| |
| num_batches = math.ceil(num_samples / args.batch_size) |
|
|
| dataloader = wds.WebLoader( |
| dataset, |
| batch_size=None, |
| shuffle=False, |
| num_workers=args.workers, |
| persistent_workers=args.workers > 0, |
| ) |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| dataloader.num_batches = num_batches |
| dataloader.num_samples = num_samples |
|
|
| return DataInfo(dataloader=dataloader, shared_epoch=shared_epoch) |
|
|
|
|
|
|
| def get_data(args, epoch=0): |
| data = {} |
|
|
| data["train"] = get_wds_dataset(args, is_train=True, epoch=epoch) |
|
|
| return data |
|
|
|
|
| if __name__ == '__main__': |
| args = parse_args(sys.argv[1:]) |
| args.workers = 10 |
| args.batch_size = 16 |
| args.world_size = 1 |
| args.num_frames = 8 |
| args.clip_type = 'vl' |
| args.model = "laion/CLIP-ViT-L-14-DataComp.XL-s13B-b90K" |
| args.train_data = '/apdcephfs_cq3/share_1311970/lb/vat2webdata/check_8frm_title_ofa_polishmplug_1tml_1dep/{00000..03020}.tar' |
| args.train_num_samples = 10_000 |
| args.dataset_type = 'webdataset' |
|
|
|
|
|
|
| data = get_data(args, epoch=0) |
|
|
| data['train'].set_epoch(0) |
| dataloader = data['train'].dataloader |
| num_batches_per_epoch = dataloader.num_batches // args.accum_freq |
| print(num_batches_per_epoch) |
|
|
|
|
| for i, batch in enumerate(tqdm(dataloader)): |
| images, input_ids, attention_mask = batch |
| |
| |