| 106 | |
| 107 | |
| 108 | class WebDataModuleFromConfig(pl.LightningDataModule): |
| 109 | def __init__(self, tar_base, batch_size, train=None, validation=None, |
| 110 | test=None, num_workers=4, multinode=True, min_size=None, |
| 111 | max_pwatermark=1.0, |
| 112 | **kwargs): |
| 113 | super().__init__(self) |
| 114 | print(f'Setting tar base to {tar_base}') |
| 115 | self.tar_base = tar_base |
| 116 | self.batch_size = batch_size |
| 117 | self.num_workers = num_workers |
| 118 | self.train = train |
| 119 | self.validation = validation |
| 120 | self.test = test |
| 121 | self.multinode = multinode |
| 122 | self.min_size = min_size # filter out very small images |
| 123 | self.max_pwatermark = max_pwatermark # filter out watermarked images |
| 124 | |
| 125 | def make_loader(self, dataset_config, train=True): |
| 126 | if 'image_transforms' in dataset_config: |
| 127 | image_transforms = [instantiate_from_config(tt) for tt in dataset_config.image_transforms] |
| 128 | else: |
| 129 | image_transforms = [] |
| 130 | |
| 131 | image_transforms.extend([torchvision.transforms.ToTensor(), |
| 132 | torchvision.transforms.Lambda(lambda x: rearrange(x * 2. - 1., 'c h w -> h w c'))]) |
| 133 | image_transforms = torchvision.transforms.Compose(image_transforms) |
| 134 | |
| 135 | if 'transforms' in dataset_config: |
| 136 | transforms_config = OmegaConf.to_container(dataset_config.transforms) |
| 137 | else: |
| 138 | transforms_config = dict() |
| 139 | |
| 140 | transform_dict = {dkey: load_partial_from_config(transforms_config[dkey]) |
| 141 | if transforms_config[dkey] != 'identity' else identity |
| 142 | for dkey in transforms_config} |
| 143 | img_key = dataset_config.get('image_key', 'jpeg') |
| 144 | transform_dict.update({img_key: image_transforms}) |
| 145 | |
| 146 | if 'postprocess' in dataset_config: |
| 147 | postprocess = instantiate_from_config(dataset_config['postprocess']) |
| 148 | else: |
| 149 | postprocess = None |
| 150 | |
| 151 | shuffle = dataset_config.get('shuffle', 0) |
| 152 | shardshuffle = shuffle > 0 |
| 153 | |
| 154 | nodesplitter = wds.shardlists.split_by_node if self.multinode else wds.shardlists.single_node_only |
| 155 | |
| 156 | if self.tar_base == "__improvedaesthetic__": |
| 157 | print("## Warning, loading the same improved aesthetic dataset " |
| 158 | "for all splits and ignoring shards parameter.") |
| 159 | tars = "pipe:aws s3 cp s3://s-laion/improved-aesthetics-laion-2B-en-subsets/aesthetics_tars/{000000..060207}.tar -" |
| 160 | else: |
| 161 | tars = os.path.join(self.tar_base, dataset_config.shards) |
| 162 | |
| 163 | dset = wds.WebDataset( |
| 164 | tars, |
| 165 | nodesplitter=nodesplitter, |