-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathpaired_transform.py
More file actions
71 lines (57 loc) · 2.03 KB
/
Copy pathpaired_transform.py
File metadata and controls
71 lines (57 loc) · 2.03 KB
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
import random
import torchvision.transforms.functional as F
from PIL import ImageFilter, Image
import torch
import numpy as np
class PairedTransform:
def __init__(
self,
size=(336, 544),
flip_prob=0.5,
rotation_degrees=20,
affine_params=None
):
self.size = size
self.flip_prob = flip_prob
self.rotation_degrees = rotation_degrees
self.affine_params = affine_params or {
"degrees": 10,
"translate": (0.05, 0.05),
"scale": (0.95, 1.05),
"shear": None
}
def __call__(self, image, mask):
# Resize
image = F.resize(image, self.size)
mask = F.resize(mask, self.size)
# Horizontal Flip
if random.random() < self.flip_prob:
image = F.hflip(image)
mask = F.hflip(mask)
# Rotation
angle = random.uniform(-self.rotation_degrees, self.rotation_degrees)
image = F.rotate(image, angle, fill=0)
mask = F.rotate(mask, angle, fill=0)
# Affine
affine_angle = random.uniform(-self.affine_params["degrees"], self.affine_params["degrees"])
translate = (
int(self.affine_params["translate"][0] * self.size[1]),
int(self.affine_params["translate"][1] * self.size[0])
)
scale = random.uniform(*self.affine_params["scale"])
shear = self.affine_params["shear"] or [0.0, 0.0]
image = F.affine(image, affine_angle, translate, scale, shear, fill=0)
mask = F.affine(mask, affine_angle, translate, scale, shear, fill=0)
# Convert to tensor
image = F.to_tensor(image)
mask = F.to_tensor(mask)
return image, mask
class DefaultPairedTransform:
def __init__(self, size=(336, 544)):
self.size = size
def __call__(self, image, mask):
image = F.resize(image, self.size)
mask = F.resize(mask, self.size)
image = F.to_tensor(image)
mask = F.to_tensor(mask)
return image, mask