Download tools/prepare_experimental_split.py from suplo6669/Enhancer: direct link, hf CLI and curl.
- Browser
- Download file 8.61 kB
-
https://huggingface.co/spaces/suplo6669/Enhancer/resolve/main/tools/prepare_experimental_split.py
- Command line
-
hf download hf://spaces/suplo6669/Enhancer/tools/prepare_experimental_split.py
-
curl -L -o prepare_experimental_split.py https://huggingface.co/spaces/suplo6669/Enhancer/resolve/main/tools/prepare_experimental_split.py
8.61 kB
| """Create a conservative experiment split without asserting identity separation.""" | |
| from __future__ import annotations | |
| import argparse | |
| from collections import Counter | |
| import copy | |
| import csv | |
| import hashlib | |
| import json | |
| from pathlib import Path | |
| import sys | |
| import cv2 | |
| import numpy as np | |
| ROOT = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(ROOT)) | |
| from tools.review_dataset import digest, safe_path | |
| def quarantine_paths(rows, groups): | |
| """Quarantine non-holdout endpoints, not identity labels or inferred groups.""" | |
| excluded, count = set(), 0 | |
| for row in rows: | |
| a, b = row['left'], row['right'] | |
| if not b: | |
| continue | |
| if a not in groups or b not in groups: | |
| raise ValueError('Review references an unknown image') | |
| if row['left_split'] != groups[a] or row['right_split'] != groups[b]: | |
| raise ValueError('Review split metadata changed') | |
| if groups[a] != groups[b]: | |
| count += 1 | |
| excluded.update(name for name in (a, b) if groups[name] != 'holdout') | |
| return excluded, count | |
| def technical_quality(image): | |
| h, w = image.shape[:2] | |
| gray = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY) | |
| scale = 256 / max(h, w) | |
| normalized = cv2.resize(gray, (max(1, round(w*scale)), max(1, round(h*scale)))) | |
| sharpness = float(cv2.Laplacian(normalized, cv2.CV_64F).var()) | |
| dark, light = float(np.mean(gray <= 5)), float(np.mean(gray >= 250)) | |
| flags = [] | |
| if min(h, w) < 256: | |
| flags.append('small_image') | |
| if sharpness < 30: | |
| flags.append('possible_global_blur') | |
| if dark > .2 or light > .2: | |
| flags.append('large_clipped_area') | |
| return {'width': w, 'height': h, 'sharpness': round(sharpness, 3), | |
| 'dark_fraction': round(dark, 4), 'light_fraction': round(light, 4), 'flags': flags} | |
| def prepare(root, source, review, output): | |
| if output.exists(): | |
| raise FileExistsError('Choose a new experiment split directory') | |
| split = json.loads((source / 'split.json').read_text(encoding='utf-8')) | |
| evidence = json.loads((review / 'review.json').read_text(encoding='utf-8')) | |
| if evidence['split_sha256'] != digest(source / 'split.json') or evidence['holdout_sha256'] != digest(source / 'holdout_paths.txt'): | |
| raise ValueError('Review no longer matches the source split') | |
| if evidence['candidates_sha256'] != digest(review / 'review_candidates.csv'): | |
| raise ValueError('Candidate evidence changed') | |
| entries = evidence['entries'] | |
| groups = {entry['path']: entry['split'] for entry in entries} | |
| if len(groups) != len(entries): | |
| raise ValueError('Repeated paths in source evidence') | |
| for group in ('train', 'validation'): | |
| names = [row['path'] for row in split[group]] | |
| if names != (source / f'{group}.txt').read_text(encoding='utf-8').splitlines(): | |
| raise ValueError('Source manifest changed') | |
| if names != [row['path'] for row in entries if row['split'] == group]: | |
| raise ValueError('Source evidence changed') | |
| with (review / 'review_candidates.csv').open(encoding='utf-8', newline='') as handle: | |
| excluded, pair_count = quarantine_paths(list(csv.DictReader(handle)), groups) | |
| reasons = {name: 'cross_split_candidate_quarantine' for name in excluded} | |
| quality, pixels = [], {} | |
| for index, entry in enumerate(entries): | |
| name = entry['path'] | |
| image = cv2.imread(str(safe_path(root, name))) | |
| if image is None: | |
| if groups[name] == 'holdout': | |
| raise ValueError('Unreadable holdout; benchmark needs repair before proceeding') | |
| excluded.add(name) | |
| reasons[name] = 'unreadable' | |
| quality.append({'path': name, 'split': groups[name], 'flags': ['unreadable']}) | |
| continue | |
| pixel_hash = hashlib.sha256(str(image.shape).encode()+image.tobytes()).hexdigest() | |
| if pixel_hash != entry['pixels_sha256']: | |
| raise ValueError(f'Image changed since review: {name}') | |
| pixels.setdefault(pixel_hash, []).append(name) | |
| stats = technical_quality(image) | |
| quality.append(dict(stats, path=name, split=groups[name])) | |
| if 'small_image' in stats['flags'] and groups[name] != 'holdout': | |
| excluded.add(name) | |
| reasons[name] = 'small_image' | |
| if (index+1) % 500 == 0: | |
| print(f'Technical checks: {index+1}/{len(entries)}', flush=True) | |
| # Preserve holdout bytes; eliminate exact-pixel overlaps and within-train duplicates. | |
| priority = {'holdout': 0, 'validation': 1, 'train': 2} | |
| for names in pixels.values(): | |
| names = sorted((n for n in names if n not in excluded), key=lambda n: (priority[groups[n]], n)) | |
| for name in names[1:]: | |
| if groups[name] != 'holdout': | |
| excluded.add(name) | |
| reasons[name] = 'exact_pixel_duplicate' | |
| result = copy.deepcopy(split) | |
| result.pop('review_receipt', None) | |
| for group in ('train', 'validation'): | |
| result[group] = [row for row in split[group] if row['path'] not in excluded] | |
| if len(result['train']) < 1000 or len(result['validation']) < 20: | |
| raise ValueError('Conservative split below experiment minimum; replenish data') | |
| result.update(review_status='experimental', lifecycle='frozen', identity_separation_verified=False, | |
| mitigation={'policy': 'quarantine_all_non_holdout_cross_split_endpoints_v1', | |
| 'source_split_sha256': digest(source/'split.json'), | |
| 'candidate_sha256': digest(review/'review_candidates.csv'), | |
| 'cross_split_rows': pair_count, 'excluded': reasons, | |
| 'technical_checks_complete': True, | |
| 'limitation': 'Unflagged identity overlap may remain; not approved for production promotion.'}) | |
| kept = {row['path'] for group in ('train', 'validation') for row in result[group]} | |
| kept.update(name for name, group in groups.items() if group == 'holdout') | |
| sample = [] | |
| for group, limit in [('train', 24), ('validation', 8), ('holdout', 8)]: | |
| rows = [r for r in quality if r['split'] == group and r['path'] in kept] | |
| rows.sort(key=lambda r: hashlib.sha256(('42:'+r['path']).encode()).hexdigest()) | |
| sample.extend(rows[:limit]) | |
| flagged = [r for r in quality if r['flags'] and r['path'] in kept] | |
| # Short inspection list, separate from automatically quarantined paths. | |
| shortlist = flagged[:12] | |
| report = {'status': 'ready_for_experiment', 'train_count': len(result['train']), | |
| 'validation_count': len(result['validation']), 'holdout_path_count': sum(g=='holdout' for g in groups.values()), | |
| 'quarantined_count': len(excluded), 'quarantine_reasons': dict(Counter(reasons.values())), | |
| 'quality_flag_count': len(flagged), 'sample': sample, 'shortlist': shortlist, | |
| 'quality': quality, 'manual_decisions_required': 0, | |
| 'identity_separation_verified': False} | |
| output.mkdir(parents=True) | |
| (output/'thumbnails').mkdir() | |
| for row in {r['path']: r for r in sample+shortlist}.values(): | |
| image = cv2.imread(str(safe_path(root, row['path']))) | |
| h,w = image.shape[:2] | |
| thumb = cv2.resize(image, (max(1,round(w*192/max(h,w))),max(1,round(h*192/max(h,w))))) | |
| name = hashlib.sha256(row['path'].encode()).hexdigest()+'.jpg' | |
| if not cv2.imwrite(str(output/'thumbnails'/name), thumb): | |
| raise OSError('Thumbnail write failed') | |
| for group in ('train', 'validation'): | |
| (output/f'{group}.txt').write_text(''.join(r['path']+'\n' for r in result[group]), encoding='utf-8') | |
| (output/'holdout_paths.txt').write_bytes((source/'holdout_paths.txt').read_bytes()) | |
| (output/'quality.json').write_text(json.dumps(report, indent=2), encoding='utf-8') | |
| result['mitigation']['quality_sha256'] = digest(output/'quality.json') | |
| (output/'split.json').write_text(json.dumps(result, indent=2), encoding='utf-8') | |
| print(json.dumps({k:v for k,v in report.items() if k not in ('sample','shortlist','quality')}), flush=True) | |
| return result, report | |
| if __name__ == '__main__': | |
| p = argparse.ArgumentParser(description=__doc__) | |
| p.add_argument('--root', type=Path, default=ROOT/'models/CodeFormer/datasets/ffhq/ffhq_512') | |
| p.add_argument('--source', type=Path, default=ROOT/'benchmarks/splits/real_portraits_v1') | |
| p.add_argument('--review', type=Path, default=ROOT/'benchmarks/reports/dataset_review_v2') | |
| p.add_argument('--output', type=Path, required=True) | |
| args = p.parse_args() | |
| prepare(args.root, args.source, args.review, args.output) | |