ema-dev / scripts /migrate_v0_to_v2.py
carlosh93's picture
migrating v0 dataset
a408130
Raw
History Blame Contribute Delete
5.05 kB
#!/usr/bin/env python3
"""Patch and validate v0 dataset JSON files for wonders_v2 compatibility."""
from __future__ import annotations
import argparse
import json
import sys
from collections import Counter
from pathlib import Path
from typing import Any
REQUIRED_FIELDS = [
"image",
"image_file",
"image_url",
"caption",
"country",
"language",
"category",
"concept",
"timestamp",
"username",
"id",
"excluded",
]
def _iter_json_files(root: Path) -> list[Path]:
files: list[Path] = []
for path in root.rglob("*.json"):
if ".cache" in path.parts:
continue
files.append(path)
return sorted(files)
def _validate_record(data: dict[str, Any]) -> list[str]:
missing = [field for field in REQUIRED_FIELDS if field not in data]
return missing
def _patch_record(
data: dict[str, Any], old_dataset: str, new_dataset: str
) -> tuple[dict[str, Any], Counter]:
stats: Counter = Counter()
if "hf_username" not in data:
data["hf_username"] = ""
stats["add_hf_username"] += 1
image_value = data.get("image")
image_file_value = data.get("image_file")
new_base = f"https://huggingface.co/datasets/{new_dataset}/resolve/main/"
if isinstance(image_file_value, str) and old_dataset in image_file_value:
data["image_file"] = image_file_value.replace(old_dataset, new_dataset)
stats["rewrite_image_file"] += 1
elif (
not isinstance(image_file_value, str)
and isinstance(image_value, str)
and image_value.strip()
):
data["image_file"] = new_base + image_value.lstrip("/")
stats["create_image_file"] += 1
return data, stats
def _write_json(path: Path, data: dict[str, Any]) -> None:
path.write_text(
json.dumps(data, indent=2, ensure_ascii=False) + "\n",
encoding="utf-8",
)
def main() -> int:
parser = argparse.ArgumentParser(
description="Migrate v0 dataset JSON files to wonders_v2-compatible format."
)
parser.add_argument(
"--source-dir",
default="./v0_data",
help="Directory containing downloaded v0 dataset files.",
)
parser.add_argument(
"--old-dataset",
required=True,
help="Old HF dataset repo id currently referenced in image_file URLs.",
)
parser.add_argument(
"--new-dataset",
required=True,
help="New HF dataset repo id to point image_file URLs to.",
)
parser.add_argument(
"--dry-run",
action="store_true",
help="Analyze and report changes without writing files.",
)
parser.add_argument(
"--max-issues",
type=int,
default=30,
help="Maximum number of validation issues to print.",
)
args = parser.parse_args()
source_dir = Path(args.source_dir).resolve()
if not source_dir.exists():
print(f"Source directory does not exist: {source_dir}")
return 2
json_files = _iter_json_files(source_dir)
if not json_files:
print(f"No JSON files found under: {source_dir}")
return 2
aggregate = Counter()
issues: list[str] = []
for path in json_files:
try:
data = json.loads(path.read_text(encoding="utf-8"))
except Exception as exc: # noqa: BLE001
issues.append(f"{path}: invalid JSON ({exc})")
continue
if not isinstance(data, dict):
issues.append(f"{path}: top-level JSON is not an object")
continue
original = json.dumps(data, sort_keys=True, ensure_ascii=False)
data, stats = _patch_record(data, args.old_dataset, args.new_dataset)
aggregate.update(stats)
missing = _validate_record(data)
if missing:
issues.append(f"{path}: missing required fields: {', '.join(missing)}")
patched = json.dumps(data, sort_keys=True, ensure_ascii=False)
if patched != original:
aggregate["changed_files"] += 1
if not args.dry_run:
_write_json(path, data)
aggregate["processed_json"] = len(json_files)
aggregate["validation_issues"] = len(issues)
print("Migration summary")
print(f"- Source dir: {source_dir}")
print(f"- Processed JSON files: {aggregate['processed_json']}")
print(f"- Changed files: {aggregate['changed_files']}")
print(f"- Added hf_username: {aggregate['add_hf_username']}")
print(f"- Rewritten image_file URLs: {aggregate['rewrite_image_file']}")
print(f"- Created missing image_file: {aggregate['create_image_file']}")
print(f"- Validation issues: {aggregate['validation_issues']}")
if issues:
print("\nValidation issues (showing up to max-issues):")
for issue in issues[: args.max_issues]:
print(f"- {issue}")
if len(issues) > args.max_issues:
print(f"... and {len(issues) - args.max_issues} more")
return 1
return 0
if __name__ == "__main__":
sys.exit(main())