disable random resize crop + ft config

This commit is contained in:
YannAhlgrim
2026-05-16 11:04:00 +02:00
parent 06f0c3fa78
commit 4df162180e
5 changed files with 63 additions and 2 deletions
@@ -0,0 +1,49 @@
meta:
model_name: vit_huge
embed_dim: 1280
load_checkpoint: true
checkpoint_folder: experiment_logs/vith14.224-bs.128-ep.300/
read_checkpoint: jepa-ep300.pth.tar
use_bfloat16: true
num_classes: 182
data:
batch_size: 128
root_path: ./wilds_data
num_workers: 10
pin_mem: true
crop_size: 224
crop_scale: [0.3, 1.0]
use_random_resized_crop: false
use_horizontal_flip: false
use_color_distortion: false
color_jitter_strength: 0.0
use_gaussian_blur: false
mask:
patch_size: 14
optimization:
optimizer: adamw # 'adamw' or 'sgd'
freeze_weights: false # true for linear probing, false for full fine-tuning
epochs: 300 # can be set higher if early_stopping
lr: 5.0e-4
weight_decay: 1.0e-2
use_cosine_schedule: true
start_lr: 0.0001
final_lr: 1.0e-06
warmup: 5
ipe_scale: 1.0
early_stopping:
enabled: true
patience: 10
min_delta: 1.0e-4
min_epochs: 15
restore_best_weights: true
validation:
eval_every: 1
logging:
folder: experiment_logs/supervised-vith14.224-bs.128-ep.300-ft/
write_tag: fine_tune
@@ -14,6 +14,7 @@ data:
pin_mem: true pin_mem: true
crop_size: 224 crop_size: 224
crop_scale: [0.3, 1.0] crop_scale: [0.3, 1.0]
use_random_resized_crop: false
use_horizontal_flip: false use_horizontal_flip: false
use_color_distortion: false use_color_distortion: false
color_jitter_strength: 0.0 color_jitter_strength: 0.0
@@ -44,5 +45,5 @@ validation:
eval_every: 1 eval_every: 1
logging: logging:
folder: experiment_logs/supervised-vith14.224-bs.128-ep.300/ folder: experiment_logs/supervised-vith14.224-bs.128-ep.300-lp/
write_tag: linear_probe write_tag: linear_probe
+2
View File
@@ -78,6 +78,7 @@ def main(args, resume_preempt=False):
use_horizontal_flip = args["data"]["use_horizontal_flip"] use_horizontal_flip = args["data"]["use_horizontal_flip"]
use_color_distortion = args["data"]["use_color_distortion"] use_color_distortion = args["data"]["use_color_distortion"]
color_jitter = args["data"]["color_jitter_strength"] color_jitter = args["data"]["color_jitter_strength"]
use_random_resized_crop = args["data"].get("use_random_resized_crop", True)
# -- # --
batch_size = args["data"]["batch_size"] batch_size = args["data"]["batch_size"]
pin_mem = args["data"]["pin_mem"] pin_mem = args["data"]["pin_mem"]
@@ -181,6 +182,7 @@ def main(args, resume_preempt=False):
horizontal_flip=use_horizontal_flip, horizontal_flip=use_horizontal_flip,
color_distortion=use_color_distortion, color_distortion=use_color_distortion,
color_jitter=color_jitter, color_jitter=color_jitter,
use_random_resized_crop=use_random_resized_crop,
) )
# -- init data-loaders/samplers # -- init data-loaders/samplers
+1
View File
@@ -239,6 +239,7 @@ def main(args, resume_preempt=False):
color_distortion=d_args["use_color_distortion"], color_distortion=d_args["use_color_distortion"],
color_jitter=d_args["color_jitter_strength"], color_jitter=d_args["color_jitter_strength"],
gaussian_blur=d_args["use_gaussian_blur"], gaussian_blur=d_args["use_gaussian_blur"],
use_random_resized_crop=d_args.get("use_random_resized_crop", True),
) )
val_transform = make_transform_eval( val_transform = make_transform_eval(
crop_size=d_args["crop_size"], crop_size=d_args["crop_size"],
+9 -1
View File
@@ -23,6 +23,7 @@ def make_transforms(
horizontal_flip=False, horizontal_flip=False,
color_distortion=False, color_distortion=False,
gaussian_blur=False, gaussian_blur=False,
use_random_resized_crop=True,
normalization=((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)), normalization=((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),
): ):
logger.info("making imagenet data transforms") logger.info("making imagenet data transforms")
@@ -36,7 +37,14 @@ def make_transforms(
return color_distort return color_distort
transform_list = [] transform_list = []
transform_list += [transforms.RandomResizedCrop(crop_size, scale=crop_scale)] if use_random_resized_crop:
transform_list += [transforms.RandomResizedCrop(crop_size, scale=crop_scale)]
else:
resize_size = int(crop_size * 256 / 224)
transform_list += [
transforms.Resize(resize_size),
transforms.CenterCrop(crop_size),
]
if horizontal_flip: if horizontal_flip:
transform_list += [transforms.RandomHorizontalFlip()] transform_list += [transforms.RandomHorizontalFlip()]
if color_distortion: if color_distortion: