data.py
2,986 bytes
| 1 | from pathlib import Path |
|---|---|
| 2 | |
| 3 | import pandas as pd |
| 4 | from PIL import Image |
| 5 | from torch.utils.data import Dataset |
| 6 | from torchvision import transforms |
| 7 | |
| 8 | IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"} |
| 9 | |
| 10 | |
| 11 | def scan_image_folder(root): |
| 12 | # Ootab paigutust root/<riik>/<pilt>. Alamkausta nimi on klassi silt. |
| 13 | root = Path(root) |
| 14 | rows = [] |
| 15 | for country_dir in sorted(p for p in root.iterdir() if p.is_dir()): |
| 16 | for img in sorted(country_dir.rglob("*")): |
| 17 | if img.suffix.lower() in IMAGE_EXTS: |
| 18 | rows.append({"path": str(img), "country": country_dir.name}) |
| 19 | if not rows: |
| 20 | raise FileNotFoundError(f"Ei leidnud ühtegi pilti kaustast {root}") |
| 21 | return pd.DataFrame(rows) |
| 22 | |
| 23 | |
| 24 | def stratified_split(df, val_frac=0.1, test_frac=0.1, seed=42): |
| 25 | # Jagame iga riigi sees eraldi, et val/test sisaldaks kõiki klasse. |
| 26 | # |
| 27 | # Väikese riigi puhul on treeningpilt tähtsam kui hindamispilt: klass, mille |
| 28 | # kohta ei ole ühtegi treeningnäidet, on mudeli väljundis olemas, aga õpitud |
| 29 | # ei ole, ja see on hullem kui puuduv mõõtmine. Seepärast jääb alati |
| 30 | # vähemalt üks pilt treeningusse ning val ja test saavad ainult selle, mis |
| 31 | # üle jääb. Suurte klasside jaotus ei muutu. |
| 32 | parts = {"train": [], "val": [], "test": []} |
| 33 | for _, group in df.groupby("country"): |
| 34 | group = group.sample(frac=1.0, random_state=seed) |
| 35 | spare = len(group) - 1 |
| 36 | n_val = min(max(1, round(len(group) * val_frac)), spare) |
| 37 | n_test = min(max(1, round(len(group) * test_frac)), spare - n_val) |
| 38 | parts["val"].append(group.iloc[:n_val]) |
| 39 | parts["test"].append(group.iloc[n_val:n_val + n_test]) |
| 40 | parts["train"].append(group.iloc[n_val + n_test:]) |
| 41 | return {name: pd.concat(chunks).reset_index(drop=True) for name, chunks in parts.items()} |
| 42 | |
| 43 | |
| 44 | def build_transforms(mean, std, image_size=224, train=True): |
| 45 | if train: |
| 46 | # Teadlikult ilma RandomHorizontalFlip'ita: liiklussuund (vasak- või |
| 47 | # parempoolne) on päris geograafiline vihje, mille peegeldamine rikuks. |
| 48 | return transforms.Compose([ |
| 49 | transforms.RandomResizedCrop(image_size, scale=(0.5, 1.0), ratio=(0.9, 1.1)), |
| 50 | transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.02), |
| 51 | transforms.ToTensor(), |
| 52 | transforms.Normalize(mean, std), |
| 53 | ]) |
| 54 | return transforms.Compose([ |
| 55 | transforms.Resize(image_size + 32), |
| 56 | transforms.CenterCrop(image_size), |
| 57 | transforms.ToTensor(), |
| 58 | transforms.Normalize(mean, std), |
| 59 | ]) |
| 60 | |
| 61 | |
| 62 | class CountryDataset(Dataset): |
| 63 | def __init__(self, df, class_to_idx, transform): |
| 64 | self.paths = df["path"].tolist() |
| 65 | self.labels = [class_to_idx[c] for c in df["country"]] |
| 66 | self.transform = transform |
| 67 | |
| 68 | def __len__(self): |
| 69 | return len(self.paths) |
| 70 | |
| 71 | def __getitem__(self, idx): |
| 72 | img = Image.open(self.paths[idx]).convert("RGB") |
| 73 | return self.transform(img), self.labels[idx] |
| 74 | |