File size: 2,346 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
"""
Shared augmentation and preprocessing transforms.
Provides train and eval transform factories based on config.
"""

from torchvision import transforms


# ImageNet normalization constants
IMAGENET_MEAN = [0.485, 0.456, 0.406]
IMAGENET_STD = [0.229, 0.224, 0.225]


def get_train_transforms(
    img_size: int = 224,
    random_crop: bool = True,
    horizontal_flip: bool = True,
    rotation_degrees: float = 15.0,
    color_jitter: bool = True,
    color_jitter_strength: float = 0.2,
) -> transforms.Compose:
    """
    Build training transforms with data augmentation.
    Avoids aggressive distortions that could alter breed-specific features.
    """
    transform_list = []

    if random_crop:
        transform_list.append(transforms.RandomResizedCrop(
            img_size,
            scale=(0.8, 1.0),
            ratio=(0.9, 1.1),
        ))
    else:
        transform_list.append(transforms.Resize((img_size, img_size)))

    if horizontal_flip:
        transform_list.append(transforms.RandomHorizontalFlip(p=0.5))

    if rotation_degrees > 0:
        transform_list.append(transforms.RandomRotation(degrees=rotation_degrees))

    if color_jitter:
        transform_list.append(transforms.ColorJitter(
            brightness=color_jitter_strength,
            contrast=color_jitter_strength,
            saturation=color_jitter_strength,
            hue=color_jitter_strength * 0.5,
        ))

    transform_list.extend([
        transforms.ToTensor(),
        transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),
    ])

    return transforms.Compose(transform_list)


def get_eval_transforms(img_size: int = 224) -> transforms.Compose:
    """
    Build evaluation/inference transforms (no augmentation).
    """
    return transforms.Compose([
        transforms.Resize((img_size, img_size)),
        transforms.ToTensor(),
        transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),
    ])


def get_denormalize_transform() -> transforms.Compose:
    """
    Inverse of ImageNet normalization for visualization.
    """
    return transforms.Compose([
        transforms.Normalize(
            mean=[0.0, 0.0, 0.0],
            std=[1.0 / s for s in IMAGENET_STD],
        ),
        transforms.Normalize(
            mean=[-m for m in IMAGENET_MEAN],
            std=[1.0, 1.0, 1.0],
        ),
    ])