File size: 3,534 Bytes
228add1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
"""
Manifest CSV dataset loader.
Loads train/val/test splits from manifest CSVs and provides PyTorch datasets.
"""

from pathlib import Path
from typing import Optional

import pandas as pd
from PIL import Image
from torch.utils.data import Dataset, DataLoader
from torchvision import transforms


class ManifestDataset(Dataset):
    """
    PyTorch Dataset that reads from a manifest CSV.
    Each row has at minimum: image_path, class_name.
    """

    def __init__(
        self,
        manifest_path: str | Path,
        data_root: str | Path,
        transform: Optional[transforms.Compose] = None,
        class_to_idx: Optional[dict[str, int]] = None,
    ):
        self.data_root = Path(data_root)
        self.transform = transform

        df = pd.read_csv(manifest_path)
        self.image_paths = df['image_path'].tolist()
        self.class_names = df['class_name'].tolist()

        # Build class mapping
        if class_to_idx is not None:
            self.class_to_idx = class_to_idx
        else:
            unique_classes = sorted(set(self.class_names))
            self.class_to_idx = {c: i for i, c in enumerate(unique_classes)}

        self.idx_to_class = {v: k for k, v in self.class_to_idx.items()}
        self.labels = [self.class_to_idx[c] for c in self.class_names]

    def __len__(self) -> int:
        return len(self.image_paths)

    def __getitem__(self, idx: int):
        img_path = self.data_root / self.image_paths[idx]
        image = Image.open(img_path).convert('RGB')

        if self.transform:
            image = self.transform(image)

        label = self.labels[idx]
        return image, label

    @property
    def num_classes(self) -> int:
        return len(self.class_to_idx)

    @property
    def classes(self) -> list[str]:
        return [self.idx_to_class[i] for i in range(self.num_classes)]


def create_dataloaders(
    manifests_dir: str | Path,
    data_root: str | Path,
    train_transform: Optional[transforms.Compose] = None,
    eval_transform: Optional[transforms.Compose] = None,
    batch_size: int = 32,
    num_workers: int = 4,
    class_to_idx: Optional[dict[str, int]] = None,
) -> dict[str, DataLoader]:
    """
    Create train/val/test dataloaders from manifest CSVs.
    Returns dict with keys 'train', 'val', 'test'.
    """
    manifests_dir = Path(manifests_dir)
    data_root = Path(data_root)

    # Build class_to_idx from training manifest to ensure consistency
    if class_to_idx is None:
        train_df = pd.read_csv(manifests_dir / 'train.csv')
        unique_classes = sorted(train_df['class_name'].unique())
        class_to_idx = {c: i for i, c in enumerate(unique_classes)}

    dataloaders = {}

    splits = {
        'train': (manifests_dir / 'train.csv', train_transform, True),
        'val': (manifests_dir / 'val.csv', eval_transform, False),
        'test': (manifests_dir / 'test.csv', eval_transform, False),
    }

    for split_name, (manifest_path, transform, shuffle) in splits.items():
        if not manifest_path.exists():
            continue

        dataset = ManifestDataset(
            manifest_path=manifest_path,
            data_root=data_root,
            transform=transform,
            class_to_idx=class_to_idx,
        )

        dataloaders[split_name] = DataLoader(
            dataset,
            batch_size=batch_size,
            shuffle=shuffle,
            num_workers=num_workers,
            pin_memory=True,
            drop_last=(split_name == 'train'),
        )

    return dataloaders