(
self,
base_path=None, metadata_path=None,
max_pixels=1920*1080, height=None, width=None,
height_division_factor=16, width_division_factor=16,
data_file_keys=("image",),
image_file_extension=("jpg", "jpeg", "png", "webp"),
repeat=1,
args=None,
)
| 12 | |
| 13 | class ImageDataset(torch.utils.data.Dataset): |
| 14 | def __init__( |
| 15 | self, |
| 16 | base_path=None, metadata_path=None, |
| 17 | max_pixels=1920*1080, height=None, width=None, |
| 18 | height_division_factor=16, width_division_factor=16, |
| 19 | data_file_keys=("image",), |
| 20 | image_file_extension=("jpg", "jpeg", "png", "webp"), |
| 21 | repeat=1, |
| 22 | args=None, |
| 23 | ): |
| 24 | if args is not None: |
| 25 | base_path = args.dataset_base_path |
| 26 | metadata_path = args.dataset_metadata_path |
| 27 | height = args.height |
| 28 | width = args.width |
| 29 | max_pixels = args.max_pixels |
| 30 | data_file_keys = args.data_file_keys.split(",") |
| 31 | repeat = args.dataset_repeat |
| 32 | |
| 33 | self.base_path = base_path |
| 34 | self.max_pixels = max_pixels |
| 35 | self.height = height |
| 36 | self.width = width |
| 37 | self.height_division_factor = height_division_factor |
| 38 | self.width_division_factor = width_division_factor |
| 39 | self.data_file_keys = data_file_keys |
| 40 | self.image_file_extension = image_file_extension |
| 41 | self.repeat = repeat |
| 42 | |
| 43 | if height is not None and width is not None: |
| 44 | print("Height and width are fixed. Setting `dynamic_resolution` to False.") |
| 45 | self.dynamic_resolution = False |
| 46 | elif height is None and width is None: |
| 47 | print("Height and width are none. Setting `dynamic_resolution` to True.") |
| 48 | self.dynamic_resolution = True |
| 49 | |
| 50 | if metadata_path is None: |
| 51 | print("No metadata. Trying to generate it.") |
| 52 | metadata = self.generate_metadata(base_path) |
| 53 | print(f"{len(metadata)} lines in metadata.") |
| 54 | self.data = [metadata.iloc[i].to_dict() for i in range(len(metadata))] |
| 55 | elif metadata_path.endswith(".json"): |
| 56 | with open(metadata_path, "r") as f: |
| 57 | metadata = json.load(f) |
| 58 | self.data = metadata |
| 59 | elif metadata_path.endswith(".jsonl"): |
| 60 | metadata = [] |
| 61 | with open(metadata_path, 'r') as f: |
| 62 | for line in tqdm(f): |
| 63 | metadata.append(json.loads(line.strip())) |
| 64 | self.data = metadata |
| 65 | else: |
| 66 | metadata = pd.read_csv(metadata_path) |
| 67 | self.data = [metadata.iloc[i].to_dict() for i in range(len(metadata))] |
| 68 | |
| 69 | |
| 70 | def generate_metadata(self, folder): |
nothing calls this directly
no test coverage detected