File size: 4,481 Bytes
c61c435
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
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