Download data.py from AI-MED-AGH/Recruitment-Task-3-Model: direct link, hf CLI and curl.
- Browser
- Download file 7.98 kB
-
https://huggingface.co/AI-MED-AGH/Recruitment-Task-3-Model/resolve/main/data.py
- Command line
-
hf download hf://AI-MED-AGH/Recruitment-Task-3-Model/data.py
-
curl -L -o data.py https://huggingface.co/AI-MED-AGH/Recruitment-Task-3-Model/resolve/main/data.py
7.98 kB
| """Immutable fold loading and in-memory preprocessing for DeepWeeds.""" | |
| from __future__ import annotations | |
| import csv | |
| import random | |
| from collections.abc import Callable, Mapping, Sequence | |
| from pathlib import Path | |
| from typing import Any | |
| import torch | |
| from PIL import Image | |
| from torch import Tensor | |
| from torch.utils.data import DataLoader, Dataset | |
| from torchvision import transforms as T | |
| REPOSITORY_ROOT = Path(__file__).resolve().parents[2] | |
| DEFAULT_LABELS_DIR = REPOSITORY_ROOT / ".artifacts" / "hf" / "Recruitment-Task-3" / "labels" | |
| DEFAULT_IMAGES_DIR = REPOSITORY_ROOT / ".artifacts" / "baseline-data" | |
| SPLIT_NAMES = ("train", "val", "test") | |
| EXPECTED_LABELS = frozenset(str(label) for label in range(9)) | |
| Sample = dict[str, str | Path] | |
| def _convert_rgb(image: Image.Image) -> Image.Image: | |
| return image.convert("RGB") | |
| def _to_unscaled_float(tensor: Tensor) -> Tensor: | |
| # Deliberately do not divide by 255. This is the controlled bad baseline. | |
| return tensor.to(dtype=torch.float32) | |
| def _good_transform() -> T.Compose: | |
| return T.Compose( | |
| [ | |
| T.Lambda(_convert_rgb), | |
| T.Resize((128, 128)), | |
| T.ToTensor(), | |
| ] | |
| ) | |
| def _bad_transform() -> T.Compose: | |
| return T.Compose( | |
| [ | |
| T.Grayscale(num_output_channels=1), | |
| T.Resize((128, 128)), | |
| T.PILToTensor(), | |
| T.Lambda(_to_unscaled_float), | |
| T.Normalize(mean=[0.5], std=[0.5]), | |
| ] | |
| ) | |
| def build_comparison_transforms() -> dict[str, T.Compose]: | |
| """Return the six fixed preprocessing pipelines for the two comparison runs.""" | |
| return { | |
| f"{run}_{split}": _good_transform() if run == "good" else _bad_transform() | |
| for run in ("good", "bad") | |
| for split in SPLIT_NAMES | |
| } | |
| def _read_split(csv_path: Path, images_dir: Path) -> list[Sample]: | |
| if not csv_path.is_file(): | |
| raise FileNotFoundError(f"Official fold CSV does not exist: {csv_path}") | |
| with csv_path.open(newline="", encoding="utf-8-sig") as csv_file: | |
| reader = csv.DictReader(csv_file) | |
| if reader.fieldnames != ["Filename", "Label"]: | |
| raise ValueError( | |
| f"Expected CSV columns ['Filename', 'Label'] in {csv_path}, " | |
| f"got {reader.fieldnames!r}" | |
| ) | |
| rows = list(reader) | |
| if not rows: | |
| raise ValueError(f"Official fold CSV is empty: {csv_path}") | |
| samples: list[Sample] = [] | |
| seen: set[str] = set() | |
| for row_number, row in enumerate(rows, start=2): | |
| filename = row["Filename"] | |
| label = row["Label"] | |
| if not filename: | |
| raise ValueError(f"Missing Filename in {csv_path} at row {row_number}") | |
| if filename in seen: | |
| raise ValueError(f"Duplicate filename in {csv_path}: {filename}") | |
| if label not in EXPECTED_LABELS: | |
| raise ValueError( | |
| f"Label must be an integer from 0 through 8 in {csv_path} " | |
| f"at row {row_number}, got {label!r}" | |
| ) | |
| image_path = images_dir / filename | |
| if not image_path.is_file(): | |
| raise FileNotFoundError(f"Image listed by {csv_path} does not exist: {image_path}") | |
| seen.add(filename) | |
| # Preserve the fold CSV label verbatim. In particular, do not replace it | |
| # from labels.csv, whose pinned upstream contents disagree for one image. | |
| samples.append({"Filename": filename, "Label": label, "path": image_path}) | |
| return samples | |
| def load_fold_splits( | |
| labels_dir: str | Path = DEFAULT_LABELS_DIR, | |
| images_dir: str | Path = DEFAULT_IMAGES_DIR, | |
| fold: int = 0, | |
| ) -> dict[str, list[Sample]]: | |
| """Load an existing official fold without inferring or rewriting membership. | |
| The values under ``Label`` remain strings exactly as published in each fold | |
| CSV. The dataset converts them to integers only when returning training items. | |
| """ | |
| if isinstance(fold, bool) or not isinstance(fold, int) or fold < 0: | |
| raise ValueError("fold must be a non-negative integer") | |
| labels_path = Path(labels_dir) | |
| images_path = Path(images_dir) | |
| if not labels_path.is_dir(): | |
| raise FileNotFoundError(f"Labels directory does not exist: {labels_path}") | |
| if not images_path.is_dir(): | |
| raise FileNotFoundError(f"Images directory does not exist: {images_path}") | |
| splits = { | |
| split: _read_split(labels_path / f"{split}_subset{fold}.csv", images_path) | |
| for split in SPLIT_NAMES | |
| } | |
| filenames = { | |
| split: {str(sample["Filename"]) for sample in samples} | |
| for split, samples in splits.items() | |
| } | |
| for index, left in enumerate(SPLIT_NAMES): | |
| for right in SPLIT_NAMES[index + 1 :]: | |
| overlap = filenames[left] & filenames[right] | |
| if overlap: | |
| example = min(overlap) | |
| raise ValueError( | |
| f"Official fold {fold} splits {left!r} and {right!r} overlap; " | |
| f"example: {example}" | |
| ) | |
| labels_csv = labels_path / "labels.csv" | |
| if labels_csv.is_file(): | |
| with labels_csv.open(newline="", encoding="utf-8-sig") as csv_file: | |
| reader = csv.DictReader(csv_file) | |
| required_columns = {"Filename", "Label"} | |
| if reader.fieldnames is None or not required_columns.issubset(reader.fieldnames): | |
| raise ValueError( | |
| f"Expected at least CSV columns ['Filename', 'Label'] in {labels_csv}, " | |
| f"got {reader.fieldnames!r}" | |
| ) | |
| population = [row["Filename"] for row in reader] | |
| if len(population) != len(set(population)): | |
| raise ValueError(f"Duplicate filename in population CSV: {labels_csv}") | |
| fold_population = set().union(*filenames.values()) | |
| if fold_population != set(population): | |
| missing = len(set(population) - fold_population) | |
| unexpected = len(fold_population - set(population)) | |
| raise ValueError( | |
| f"Official fold {fold} does not cover labels.csv exactly " | |
| f"(missing={missing}, unexpected={unexpected})" | |
| ) | |
| return splits | |
| class DeepWeedsDataset(Dataset[tuple[Tensor, int, str]]): | |
| """Read source images lazily and apply one selected transform in memory.""" | |
| def __init__( | |
| self, | |
| samples: Sequence[Mapping[str, Any]], | |
| transform: Callable[[Image.Image], Tensor], | |
| ) -> None: | |
| if transform is None: | |
| raise ValueError("transform is required") | |
| self.samples = list(samples) | |
| self.transform = transform | |
| def __len__(self) -> int: | |
| return len(self.samples) | |
| def __getitem__(self, index: int) -> tuple[Tensor, int, str]: | |
| sample = self.samples[index] | |
| image_path = Path(sample["path"]) | |
| with Image.open(image_path) as image: | |
| image.load() | |
| tensor = self.transform(image) | |
| return tensor, int(sample["Label"]), str(sample["Filename"]) | |
| def _seed_worker(_worker_id: int) -> None: | |
| worker_seed = torch.initial_seed() % (2**32) | |
| random.seed(worker_seed) | |
| def make_loader( | |
| samples: Sequence[Mapping[str, Any]], | |
| transform: Callable[[Image.Image], Tensor], | |
| shuffle: bool, | |
| seed: int, | |
| batch_size: int = 64, | |
| num_workers: int = 0, | |
| pin_memory: bool = False, | |
| ) -> DataLoader[tuple[Tensor, int, str]]: | |
| """Build a seeded loader; callers explicitly choose whether it shuffles.""" | |
| if batch_size <= 0: | |
| raise ValueError("batch_size must be positive") | |
| if num_workers < 0: | |
| raise ValueError("num_workers must be non-negative") | |
| generator = torch.Generator() | |
| generator.manual_seed(seed) | |
| return DataLoader( | |
| DeepWeedsDataset(samples, transform), | |
| batch_size=batch_size, | |
| shuffle=shuffle, | |
| num_workers=num_workers, | |
| pin_memory=pin_memory, | |
| generator=generator, | |
| worker_init_fn=_seed_worker, | |
| ) | |