Piotrek-max's picture
Publish reproducible DeepWeeds recruitment baseline
301f30e verified
Raw History Blame Contribute Delete
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,
)