dataset-utils / resize /preprocess_dataset.py
recoilme's picture
Upload folder using huggingface_hub
c9a5b52 verified
Raw
History Blame Contribute Delete
5.34 kB
"""Pre-process a dataset for model training: resize & center-crop every image
into a bucket within [min, max] on the long side, aligned to a step.
Images are written as sequentially numbered files with paired .txt captions.
Preserves original format (JPG q97, PNG lossless).
Usage:
python preprocess_dataset.py --input /path/to/img --output /path/to/out
python preprocess_dataset.py --input /path/to/img --output /path/to/out --min-size 1024 --max-size 1152 --step 64
python preprocess_dataset.py --input /path/to/img --output /path/to/out --dry-run
"""
import argparse
import os
import struct
from PIL import Image, ImageFile
ImageFile.LOAD_TRUNCATED_IMAGES = True
from tqdm import tqdm
JPEG_QUALITY = 97
NAME_WIDTH = 7
IMAGE_EXTS = {'.jpg', '.jpeg', '.png', '.webp', '.bmp'}
def process_image(img, min_size, max_size, step):
"""Resize + crop to fit [min_size, max_size] with step alignment."""
w, h = img.size
if max(w, h) > max_size:
if w >= h:
new_w, new_h = max_size, int(max_size * h / w)
else:
new_w, new_h = int(max_size * w / h), max_size
else:
new_w, new_h = w, h
if min(new_w, new_h) < min_size:
if new_w <= new_h:
new_w, new_h = min_size, int(min_size * new_h / new_w)
else:
new_w, new_h = int(min_size * new_w / new_h), min_size
crop_w = min(max_size, (new_w // step) * step)
crop_h = min(max_size, (new_h // step) * step)
crop_w = max(min_size, crop_w)
crop_h = max(min_size, crop_h)
img = img.convert("RGB").resize((new_w, new_h), Image.LANCZOS)
left = (new_w - crop_w) // 2
top = (new_h - crop_h) // 2
return img.crop((left, top, left + crop_w, top + crop_h))
def main():
parser = argparse.ArgumentParser(description="Resize dataset into buckets")
parser.add_argument("--input", required=True)
parser.add_argument("--output", required=True)
parser.add_argument("--min-size", type=int, default=768, help="bucket min side")
parser.add_argument("--max-size", type=int, default=1280, help="bucket max side")
parser.add_argument("--step", type=int, default=16, help="bucket alignment step")
parser.add_argument("--dry-run", action="store_true")
args = parser.parse_args()
os.makedirs(args.output, exist_ok=True)
# Recursive scan
print("Scanning recursively...")
captions = {} # full_path_without_ext → text
image_paths = [] # full paths
for root, dirs, files in os.walk(args.input):
txt_stems = set()
for fname in files:
if fname.lower().endswith('.txt'):
stem = os.path.join(root, os.path.splitext(fname)[0])
txt_stems.add(stem)
with open(os.path.join(root, fname)) as f:
captions[stem] = f.read()
for fname in sorted(files):
if os.path.splitext(fname)[1].lower() in IMAGE_EXTS:
stem = os.path.join(root, os.path.splitext(fname)[0])
if stem in txt_stems:
image_paths.append(os.path.join(root, fname))
image_paths.sort()
print(f"Found {len(image_paths)} images, {len(captions)} paired .txt captions")
print(f"Min: {args.min_size}, Max: {args.max_size}, Step: {args.step}, Quality: {JPEG_QUALITY}")
if args.dry_run:
for path in tqdm(image_paths, desc="Dry run"):
try:
img = Image.open(path)
img.load()
out = process_image(img, args.min_size, args.max_size, args.step)
print(f"{path}: {img.size} -> {out.size}")
except (OSError, IOError, struct.error) as e:
print(f"{path}: SKIP ({e})")
return
stats = {"upscaled": 0, "downscaled": 0, "unchanged": 0, "skipped": 0}
for idx, src_path in enumerate(tqdm(image_paths, desc="Processing"), start=1):
stem_no_ext = os.path.splitext(src_path)[0]
ext = os.path.splitext(src_path)[1].lower()
try:
img = Image.open(src_path)
img.load()
orig_w, orig_h = img.size
out = process_image(img, args.min_size, args.max_size, args.step)
except (OSError, IOError, struct.error) as e:
tqdm.write(f" Skip {os.path.basename(src_path)}: {e}")
stats["skipped"] += 1
continue
if out.size == (orig_w, orig_h):
stats["unchanged"] += 1
elif max(out.size) > max(orig_w, orig_h):
stats["upscaled"] += 1
else:
stats["downscaled"] += 1
# Sequential names + preserve format
name = f"{idx:0{NAME_WIDTH}d}"
out_path = os.path.join(args.output, f"{name}{ext}")
if ext in {'.jpg', '.jpeg'}:
out.save(out_path, quality=JPEG_QUALITY)
else:
out.save(out_path)
# Copy paired .txt
if stem_no_ext in captions:
txt_path = os.path.join(args.output, f"{name}.txt")
with open(txt_path, "w") as f:
f.write(captions[stem_no_ext])
print(f"\nDone: {stats['unchanged']} unchanged, "
f"{stats['upscaled']} upscaled, "
f"{stats['downscaled']} downscaled, "
f"{stats['skipped']} skipped")
print(f"Output: {args.output}")
if __name__ == "__main__":
main()