
N : 적용할 augmentation 개수.
M : 적용한 augmentation의 변환 정도.
def augment_list():
l = [
(AutoContrast, 0, 1),
(Equalize, 0, 1),
(Invert, 0, 1),
(Rotate, 0, 30),
(Posterize, 0, 4),
(Solarize, 0, 256),
(SolarizeAdd, 0, 110),
(Color, 0.1, 1.9),
(Contrast, 0.1, 1.9),
(Brightness, 0.1, 1.9),
(Sharpness, 0.1, 1.9),
(ShearX, 0., 0.3),
(ShearY, 0., 0.3),
(CutoutAbs, 0, 40),
(TranslateXabs, 0., 100),
(TranslateYabs, 0., 100),
]
return l
class RandAugment:
def __init__(self, n, m):
self.n = n
self.m = m # [0, 30]
self.augment_list = augment_list()
def __call__(self, img):
ops = random.choices(self.augment_list, k=self.n)
org_img = img.copy()
for op, minval, maxval in ops:
val = (float(self.m) / 30) * float(maxval - minval) + minval
img_list.append(op(org_img, val))
fn_names.append(str(op).split(' ')[1])
img = op(img, val)
return img
if __name__ == '__main__':
import matplotlib.pyplot as plt
def visualize(images, names):
fig = plt.figure(figsize=(10, 10))
for i, (img, name) in enumerate(zip(images, names)):
fig.add_subplot(3, 3, i+1)
plt.imshow(img)
plt.title(name)
plt.show()
path = '../dogs/Golden retriever/n158409.jpeg'
image = Image.open(path).convert('RGB')
img_list = [image]
fn_names = ['Original']
ra = RandAugment(3,2)
transform_img = ra(image)
img_list.append(transform_img)
fn_names.append('All')
visualize(img_list, fn_names)

N=3, M=2 일 때.