from __future__ import annotations import hashlib from dataclasses import asdict, dataclass, field from pathlib import Path IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"} VIDEO_EXTENSIONS = {".mp4", ".mov", ".mkv", ".webm", ".avi"} TEXT_EXTENSIONS = {".txt", ".caption", ".jsonl", ".json"} @dataclass(slots=True) class DatasetItem: path: str kind: str size_bytes: int width: int = 0 height: int = 0 caption_path: str = "" duplicate_key: str = "" @dataclass(slots=True) class DatasetReport: path: str total_files: int = 0 image_count: int = 0 video_count: int = 0 text_count: int = 0 caption_count: int = 0 missing_caption_count: int = 0 duplicate_groups: int = 0 dimensions: dict[str, int] = field(default_factory=dict) extensions: dict[str, int] = field(default_factory=dict) items: list[DatasetItem] = field(default_factory=list) warnings: list[str] = field(default_factory=list) def to_dict(self) -> dict: return asdict(self) def _hash_file(path: Path) -> str: digest = hashlib.sha1() with path.open("rb") as handle: for chunk in iter(lambda: handle.read(1024 * 1024), b""): digest.update(chunk) return digest.hexdigest() def scan_dataset(path: str | Path, *, limit: int = 500) -> DatasetReport: root = Path(path).expanduser().resolve() report = DatasetReport(path=str(root)) if not root.is_dir(): report.warnings.append("Dataset folder does not exist.") return report hashes: dict[str, int] = {} try: files = [item for item in root.rglob("*") if item.is_file()] except OSError as exc: report.warnings.append(f"Dataset could not be scanned: {exc}") return report report.total_files = len(files) for item in files: suffix = item.suffix.casefold() report.extensions[suffix or "(none)"] = report.extensions.get(suffix or "(none)", 0) + 1 if suffix in IMAGE_EXTENSIONS: report.image_count += 1 if not any(item.with_suffix(ext).is_file() for ext in (".txt", ".caption")): report.missing_caption_count += 1 elif suffix in VIDEO_EXTENSIONS: report.video_count += 1 if suffix in TEXT_EXTENSIONS: report.text_count += 1 if suffix in {".txt", ".caption"}: report.caption_count += 1 for item in files[: max(1, limit)]: suffix = item.suffix.casefold() kind = "other" width = height = 0 caption_path = "" duplicate_key = "" if suffix in IMAGE_EXTENSIONS: kind = "image" caption = next((item.with_suffix(ext) for ext in (".txt", ".caption") if item.with_suffix(ext).is_file()), None) caption_path = str(caption) if caption else "" try: from PIL import Image with Image.open(item) as image: width, height = image.size label = f"{width}x{height}" report.dimensions[label] = report.dimensions.get(label, 0) + 1 except Exception: pass try: duplicate_key = _hash_file(item) hashes[duplicate_key] = hashes.get(duplicate_key, 0) + 1 except OSError: duplicate_key = "" elif suffix in VIDEO_EXTENSIONS: kind = "video" elif suffix in TEXT_EXTENSIONS: kind = "text" try: size = item.stat().st_size except OSError: size = 0 report.items.append( DatasetItem( path=str(item), kind=kind, size_bytes=size, width=width, height=height, caption_path=caption_path, duplicate_key=duplicate_key, ) ) if report.total_files > limit: report.warnings.append(f"Showing first {limit:,} files; totals still include all files.") report.duplicate_groups = sum(1 for count in hashes.values() if count > 1) if report.image_count and report.missing_caption_count: report.warnings.append(f"{report.missing_caption_count:,} sampled image(s) do not have sidecar captions.") if report.duplicate_groups: report.warnings.append(f"{report.duplicate_groups:,} duplicate image group(s) found in the sample.") return report