File size: 5,071 Bytes
69ce2b7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
import torch
import torchvision.transforms.functional as TF
from torchvision import transforms
from typing import Tuple, List
import random


class Compose:
    """Apply a sequence of transforms to both image and bbox."""
    def __init__(self, transforms_list):
        self.transforms = transforms_list

    def __call__(self, image, bbox):
        for t in self.transforms:
            image, bbox = t(image, bbox)
        return image, bbox


class ToTensor:
    """Convert PIL image to tensor. Bbox is already a list, pass through."""
    def __call__(self, image, bbox):
        return TF.to_tensor(image), bbox


class Normalize:
    """Normalize image tensor. Bbox unchanged."""
    def __init__(self, mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)):
        self.mean = mean
        self.std = std

    def __call__(self, image, bbox):
        return TF.normalize(image, self.mean, self.std), bbox


class Resize:
    """
    Resize image to (size, size).
    Bbox is normalized [0,1] so no change needed.
    """
    def __init__(self, size: int = 224):
        self.size = size

    def __call__(self, image, bbox):
        image = TF.resize(image, [self.size, self.size])
        return image, bbox


class RandomHorizontalFlip:
    """
    Flip image horizontally with probability p.
    Bbox x-coords must be mirrored: x_min' = 1 - x_max, x_max' = 1 - x_min.
    """
    def __init__(self, p: float = 0.5):
        self.p = p

    def __call__(self, image, bbox):
        if random.random() < self.p:
            image = TF.hflip(image)
            x_min, y_min, x_max, y_max = bbox
            bbox = [1.0 - x_max, y_min, 1.0 - x_min, y_max]
        return image, bbox


class RandomColorJitter:
    """Color jitter on image only. Bbox unchanged."""
    def __init__(self, brightness=0.3, contrast=0.3, saturation=0.3, hue=0.1):
        self.jitter = transforms.ColorJitter(
            brightness=brightness,
            contrast=contrast,
            saturation=saturation,
            hue=hue,
        )

    def __call__(self, image, bbox):
        return self.jitter(image), bbox


class RandomCrop:
    """
    Crop a random region of the image, keeping at least min_overlap
    of the primary bbox inside the crop. Falls back to center crop
    if a valid crop isn't found within max_attempts.

    Bbox is recalculated relative to the cropped region.
    """
    def __init__(self, size: int = 224, min_overlap: float = 0.7, max_attempts: int = 10):
        self.size = size
        self.min_overlap = min_overlap
        self.max_attempts = max_attempts

    def __call__(self, image, bbox):
        w, h = image.size  # PIL: (width, height)
        x_min, y_min, x_max, y_max = bbox

        # convert normalized bbox to pixel coords
        bx1, by1 = x_min * w, y_min * h
        bx2, by2 = x_max * w, y_max * h

        crop_w = min(self.size, w)
        crop_h = min(self.size, h)

        for _ in range(self.max_attempts):
            left = random.randint(0, max(0, w - crop_w))
            top  = random.randint(0, max(0, h - crop_h))
            right  = left + crop_w
            bottom = top  + crop_h

            # intersection with bbox
            ix1 = max(bx1, left)
            iy1 = max(by1, top)
            ix2 = min(bx2, right)
            iy2 = min(by2, bottom)

            if ix2 > ix1 and iy2 > iy1:
                inter_area = (ix2 - ix1) * (iy2 - iy1)
                bbox_area  = (bx2 - bx1) * (by2 - by1)
                if bbox_area > 0 and (inter_area / bbox_area) >= self.min_overlap:
                    image = TF.crop(image, top, left, crop_h, crop_w)
                    image = TF.resize(image, [self.size, self.size])

                    # recalculate bbox relative to crop, re-normalize
                    new_bbox = [
                        (bx1 - left) / crop_w,
                        (by1 - top)  / crop_h,
                        (bx2 - left) / crop_w,
                        (by2 - top)  / crop_h,
                    ]
                    new_bbox = [max(0.0, min(1.0, v)) for v in new_bbox]
                    return image, new_bbox

        # fallback: center crop
        left = (w - crop_w) // 2
        top  = (h - crop_h) // 2
        image = TF.crop(image, top, left, crop_h, crop_w)
        image = TF.resize(image, [self.size, self.size])
        new_bbox = [
            (bx1 - left) / crop_w,
            (by1 - top)  / crop_h,
            (bx2 - left) / crop_w,
            (by2 - top)  / crop_h,
        ]
        new_bbox = [max(0.0, min(1.0, v)) for v in new_bbox]
        return image, new_bbox


def get_train_transforms(size: int = 224) -> Compose:
    return Compose([
        Resize(size + 32),           # resize slightly larger first
        RandomCrop(size),            # then random crop to target size
        RandomHorizontalFlip(p=0.5),
        RandomColorJitter(),
        ToTensor(),
        Normalize(),
    ])


def get_val_transforms(size: int = 224) -> Compose:
    return Compose([
        Resize(size),
        ToTensor(),
        Normalize(),
    ])