clean.py
4,037 bytes
| 1 | import argparse |
|---|---|
| 2 | from pathlib import Path |
| 3 | |
| 4 | import imagehash |
| 5 | import numpy as np |
| 6 | import pandas as pd |
| 7 | from PIL import Image |
| 8 | from tqdm import tqdm |
| 9 | |
| 10 | from .data import scan_image_folder |
| 11 | |
| 12 | |
| 13 | def sharpness_score(gray): |
| 14 | # Gradiendi energia keskmine: udusel pildil on see madal. |
| 15 | gy, gx = np.gradient(gray) |
| 16 | return float(np.mean(gx * gx + gy * gy)) |
| 17 | |
| 18 | |
| 19 | def inspect_image(path, min_side, dark_limit, bright_limit, min_sharpness): |
| 20 | # Tagastab (phash, None) kui pilt kõlbab, muidu (None, põhjus). |
| 21 | try: |
| 22 | with Image.open(path) as im: |
| 23 | im = im.convert("RGB") |
| 24 | except Exception: |
| 25 | return None, "rikutud fail" |
| 26 | if min(im.size) < min_side: |
| 27 | return None, "liiga väike" |
| 28 | small = np.asarray(im.convert("L").resize((128, 128)), dtype=np.float32) |
| 29 | mean = small.mean() |
| 30 | if mean < dark_limit: |
| 31 | return None, "liiga tume" |
| 32 | if mean > bright_limit: |
| 33 | return None, "liiga hele" |
| 34 | if sharpness_score(small) < min_sharpness: |
| 35 | return None, "liiga udune" |
| 36 | return imagehash.phash(im), None |
| 37 | |
| 38 | |
| 39 | def main(): |
| 40 | ap = argparse.ArgumentParser(description="Puhastab toorandmestiku ja kirjutab manifesti") |
| 41 | ap.add_argument("--raw-dir", required=True, help="kaust, kus iga alamkaust on üks riik") |
| 42 | ap.add_argument("--out-manifest", default="data/manifest.csv") |
| 43 | ap.add_argument("--report", default="data/cleaning_report.csv") |
| 44 | ap.add_argument("--min-per-class", type=int, default=100) |
| 45 | ap.add_argument("--max-per-class", type=int, default=2000) |
| 46 | ap.add_argument("--min-side", type=int, default=128) |
| 47 | ap.add_argument("--dark-limit", type=float, default=20.0) |
| 48 | ap.add_argument("--bright-limit", type=float, default=235.0) |
| 49 | ap.add_argument("--min-sharpness", type=float, default=25.0) |
| 50 | ap.add_argument("--seed", type=int, default=42) |
| 51 | args = ap.parse_args() |
| 52 | |
| 53 | df = scan_image_folder(args.raw_dir) |
| 54 | print(f"Leitud {len(df)} pilti, {df['country'].nunique()} riiki") |
| 55 | |
| 56 | kept, dropped = [], [] |
| 57 | seen_hashes = {} |
| 58 | for row in tqdm(df.itertuples(), total=len(df), desc="Kontrollin pilte"): |
| 59 | phash, reason = inspect_image( |
| 60 | row.path, args.min_side, args.dark_limit, args.bright_limit, args.min_sharpness |
| 61 | ) |
| 62 | if reason is None: |
| 63 | key = str(phash) |
| 64 | if key in seen_hashes: |
| 65 | reason = f"duplikaat ({seen_hashes[key]})" |
| 66 | else: |
| 67 | seen_hashes[key] = row.path |
| 68 | kept.append({"path": row.path, "country": row.country}) |
| 69 | if reason is not None: |
| 70 | dropped.append({"path": row.path, "country": row.country, "reason": reason}) |
| 71 | |
| 72 | kept_df = pd.DataFrame(kept) |
| 73 | |
| 74 | # Klasside piirid: liiga väikesed riigid välja, liiga suured lakke. |
| 75 | counts = kept_df["country"].value_counts() |
| 76 | small = counts[counts < args.min_per_class].index |
| 77 | for country in small: |
| 78 | rows = kept_df[kept_df["country"] == country] |
| 79 | dropped.extend( |
| 80 | {"path": p, "country": country, "reason": "riigil liiga vähe pilte"} |
| 81 | for p in rows["path"] |
| 82 | ) |
| 83 | kept_df = kept_df[~kept_df["country"].isin(small)] |
| 84 | capped = [ |
| 85 | g.sample(min(len(g), args.max_per_class), random_state=args.seed) |
| 86 | for _, g in kept_df.groupby("country") |
| 87 | ] |
| 88 | kept_df = pd.concat(capped).reset_index(drop=True) |
| 89 | |
| 90 | Path(args.out_manifest).parent.mkdir(parents=True, exist_ok=True) |
| 91 | kept_df.to_csv(args.out_manifest, index=False) |
| 92 | Path(args.report).parent.mkdir(parents=True, exist_ok=True) |
| 93 | dropped_df = pd.DataFrame(dropped, columns=["path", "country", "reason"]) |
| 94 | dropped_df.to_csv(args.report, index=False) |
| 95 | |
| 96 | print(f"Alles {len(kept_df)} pilti, {kept_df['country'].nunique()} riiki") |
| 97 | if len(dropped_df): |
| 98 | # Duplikaatide põhjus sisaldab faili teed, koondame üldnime alla. |
| 99 | reasons = dropped_df["reason"].str.replace(r"duplikaat \(.*\)", "duplikaat", regex=True) |
| 100 | print("Välja jäetud põhjuste kaupa:") |
| 101 | print(reasons.value_counts().to_string()) |
| 102 | |
| 103 | |
| 104 | if __name__ == "__main__": |
| 105 | main() |
| 106 | |