add evaluation transform method

This commit is contained in:
YannAhlgrim
2026-04-25 11:01:38 +02:00
parent 02991a7eb5
commit c54816aa93
+21 -7
View File
@@ -23,19 +23,16 @@ def make_transforms(
horizontal_flip=False, horizontal_flip=False,
color_distortion=False, color_distortion=False,
gaussian_blur=False, gaussian_blur=False,
normalization=((0.485, 0.456, 0.406), normalization=((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),
(0.229, 0.224, 0.225))
): ):
logger.info('making imagenet data transforms') logger.info("making imagenet data transforms")
def get_color_distortion(s=1.0): def get_color_distortion(s=1.0):
# s is the strength of color distortion. # s is the strength of color distortion.
color_jitter = transforms.ColorJitter(0.8 * s, 0.8 * s, 0.8 * s, 0.2 * s) color_jitter = transforms.ColorJitter(0.8 * s, 0.8 * s, 0.8 * s, 0.2 * s)
rnd_color_jitter = transforms.RandomApply([color_jitter], p=0.8) rnd_color_jitter = transforms.RandomApply([color_jitter], p=0.8)
rnd_gray = transforms.RandomGrayscale(p=0.2) rnd_gray = transforms.RandomGrayscale(p=0.2)
color_distort = transforms.Compose([ color_distort = transforms.Compose([rnd_color_jitter, rnd_gray])
rnd_color_jitter,
rnd_gray])
return color_distort return color_distort
transform_list = [] transform_list = []
@@ -53,8 +50,25 @@ def make_transforms(
return transform return transform
def make_transform_eval(
crop_size=224, normalization=((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))
):
logger.info("making imagenet evaluation transforms")
resize_size = int(crop_size * 256 / 224)
transform = transforms.Compose(
[
transforms.Resize(resize_size),
transforms.CenterCrop(crop_size),
transforms.ToTensor(),
transforms.Normalize(normalization[0], normalization[1]),
]
)
return transform
class GaussianBlur(object): class GaussianBlur(object):
def __init__(self, p=0.5, radius_min=0.1, radius_max=2.): def __init__(self, p=0.5, radius_min=0.1, radius_max=2.0):
self.prob = p self.prob = p
self.radius_min = radius_min self.radius_min = radius_min
self.radius_max = radius_max self.radius_max = radius_max