| 52 | |
| 53 | |
| 54 | class WebDataModuleFromConfig(pl.LightningDataModule): |
| 55 | def __init__(self, |
| 56 | tar_base, # can be a list of paths or a single path |
| 57 | batch_size, |
| 58 | val_batch_size=None, |
| 59 | train=None, |
| 60 | validation=None, |
| 61 | test=None, |
| 62 | num_workers=4, |
| 63 | val_num_workers: int = None, |
| 64 | multinode=True, |
| 65 | remove_keys: list = None, # list of keys to remove from the sample |
| 66 | ): |
| 67 | super().__init__() |
| 68 | if isinstance(tar_base, str): |
| 69 | self.tar_base = tar_base |
| 70 | elif isinstance(tar_base, ListConfig) or isinstance(tar_base, list): |
| 71 | # check which tar_base exists |
| 72 | for path in tar_base: |
| 73 | if os.path.exists(path): |
| 74 | self.tar_base = path |
| 75 | break |
| 76 | else: |
| 77 | raise FileNotFoundError("Could not find a valid tarbase.") |
| 78 | else: |
| 79 | raise ValueError(f'Invalid tar_base type {type(tar_base)}') |
| 80 | print(f'[WebDataModuleFromConfig] Setting tar base to {self.tar_base}') |
| 81 | |
| 82 | self.batch_size = batch_size |
| 83 | self.num_workers = num_workers |
| 84 | self.train = train |
| 85 | self.validation = validation |
| 86 | self.test = test |
| 87 | self.multinode = multinode |
| 88 | self.val_batch_size = val_batch_size if val_batch_size is not None else batch_size |
| 89 | self.val_num_workers = val_num_workers if val_num_workers is not None else num_workers |
| 90 | self.rm_keys = remove_keys if remove_keys is not None else [] |
| 91 | |
| 92 | def make_loader(self, dataset_config, train=True): |
| 93 | image_transforms = [] |
| 94 | lambda_fn = lambda x: x * 2. - 1. # normalize to [-1, 1] |
| 95 | image_transforms.extend([torchvision.transforms.ToTensor(), |
| 96 | torchvision.transforms.Lambda(lambda_fn)]) |
| 97 | if 'image_transforms' in dataset_config: |
| 98 | image_transforms.extend([instantiate_from_config(tt) for tt in dataset_config.image_transforms]) |
| 99 | image_transforms = torchvision.transforms.Compose(image_transforms) |
| 100 | |
| 101 | if 'transforms' in dataset_config: |
| 102 | transforms_config = OmegaConf.to_container(dataset_config.transforms) |
| 103 | else: |
| 104 | transforms_config = dict() |
| 105 | |
| 106 | transform_dict = {dkey: load_partial_from_config(transforms_config[dkey]) |
| 107 | if transforms_config[dkey] != 'identity' else identity |
| 108 | for dkey in transforms_config} |
| 109 | # this is crucial to set correct image key to get the transofrms applied correctly |
| 110 | img_keys = dataset_config.get('image_key', 'image.png') |
| 111 | if isinstance(img_keys, str): |
nothing calls this directly
no outgoing calls
no test coverage detected