MCPcopy Create free account
hub / github.com/MegaScenes/nvs / WebDataModuleFromConfig

Class WebDataModuleFromConfig

ldm/data/laion.py:108–218  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

106
107
108class 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,

Callers 1

example02Function · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected