Untitled

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)

Untitled

N=3, M=2 일 때.