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