Spaces:
Paused
Paused
Download ldm/data/base.py from NguyenDinhHieu/EquiFashion: direct link, hf CLI and curl.
- Browser
- Download file 1.2 kB
-
https://huggingface.co/spaces/NguyenDinhHieu/EquiFashion/resolve/c9aa3dbc14d03d7433dc65376c547e34a8ec8ab9/ldm/data/base.py
- Command line
-
hf download hf://spaces/NguyenDinhHieu/EquiFashion@c9aa3dbc14d03d7433dc65376c547e34a8ec8ab9/ldm/data/base.py
-
curl -L -o base.py https://huggingface.co/spaces/NguyenDinhHieu/EquiFashion/resolve/c9aa3dbc14d03d7433dc65376c547e34a8ec8ab9/ldm/data/base.py
1.2 kB
| import os | |
| import numpy as np | |
| from abc import abstractmethod | |
| from torch.utils.data import Dataset, ConcatDataset, ChainDataset, IterableDataset | |
| class Txt2ImgIterableBaseDataset(IterableDataset): | |
| ''' | |
| Define an interface to make the IterableDatasets for text2img data chainable | |
| ''' | |
| def __init__(self, num_records=0, valid_ids=None, size=256): | |
| super().__init__() | |
| self.num_records = num_records | |
| self.valid_ids = valid_ids | |
| self.sample_ids = valid_ids | |
| self.size = size | |
| print(f'{self.__class__.__name__} dataset contains {self.__len__()} examples.') | |
| def __len__(self): | |
| return self.num_records | |
| def __iter__(self): | |
| pass | |
| class PRNGMixin(object): | |
| """ | |
| Adds a prng property which is a numpy RandomState which gets | |
| reinitialized whenever the pid changes to avoid synchronized sampling | |
| behavior when used in conjunction with multiprocessing. | |
| """ | |
| def prng(self): | |
| currentpid = os.getpid() | |
| if getattr(self, "_initpid", None) != currentpid: | |
| self._initpid = currentpid | |
| self._prng = np.random.RandomState() | |
| return self._prng |