SyntheticMDProductions's picture
Update ADAM safety, UI, and model workflows (#1)
c61c435
Raw History Blame Contribute Delete
4.48 kB
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