From 4d2fd1021b2dbd297f5a26501129ace70304de84 Mon Sep 17 00:00:00 2001 From: YannAhlgrim Date: Thu, 2 Apr 2026 14:15:06 +0200 Subject: [PATCH] integrate make_iwildcam in train.py --- src/datasets/wilds.py | 4 +- src/train.py | 288 ++++++++++++++++++++++-------------------- 2 files changed, 156 insertions(+), 136 deletions(-) diff --git a/src/datasets/wilds.py b/src/datasets/wilds.py index 4f830f1..c6983ba 100644 --- a/src/datasets/wilds.py +++ b/src/datasets/wilds.py @@ -8,6 +8,7 @@ logger = getLogger() def make_iwildcam( transform, batch_size, + collator=None, split="extra_unlabeled", num_workers=8, world_size=1, @@ -36,12 +37,13 @@ def make_iwildcam( else: dataset = WildsToTorchWrapper(dataset) - dist_sampler = torch.utils.data.DistributedSampler( + dist_sampler = torch.utils.data.distributed.DistributedSampler( dataset=dataset, num_replicas=world_size, rank=rank, shuffle=shuffle ) data_loader = torch.utils.data.DataLoader( dataset, + collate_fn=collator, sampler=dist_sampler, batch_size=batch_size, drop_last=drop_last, diff --git a/src/train.py b/src/train.py index 0f387d6..c79a735 100644 --- a/src/train.py +++ b/src/train.py @@ -13,7 +13,7 @@ try: # -- SURE TO UPDATE THIS TO GET LOCAL-RANK ON NODE, OR ENSURE # -- THAT YOUR JOBS ARE LAUNCHED WITH ONLY 1 DEVICE VISIBLE # -- TO EACH PROCESS - os.environ['CUDA_VISIBLE_DEVICES'] = os.environ['SLURM_LOCALID'] + os.environ["CUDA_VISIBLE_DEVICES"] = os.environ["SLURM_LOCALID"] except Exception: pass @@ -31,22 +31,12 @@ from torch.nn.parallel import DistributedDataParallel from src.masks.multiblock import MaskCollator as MBMaskCollator from src.masks.utils import apply_masks -from src.utils.distributed import ( - init_distributed, - AllReduce -) -from src.utils.logging import ( - CSVLogger, - gpu_timer, - grad_logger, - AverageMeter) +from src.utils.distributed import init_distributed, AllReduce +from src.utils.logging import CSVLogger, gpu_timer, grad_logger, AverageMeter from src.utils.tensors import repeat_interleave_batch -from src.datasets.imagenet1k import make_imagenet1k +from src.datasets.wilds import make_iwildcam -from src.helper import ( - load_checkpoint, - init_model, - init_opt) +from src.helper import load_checkpoint, init_model, init_opt from src.transforms import make_transforms # -- @@ -65,98 +55,101 @@ logger = logging.getLogger() def main(args, resume_preempt=False): - # ----------------------------------------------------------------------- # # PASSED IN PARAMS FROM CONFIG FILE # ----------------------------------------------------------------------- # # -- META - use_bfloat16 = args['meta']['use_bfloat16'] - model_name = args['meta']['model_name'] - load_model = args['meta']['load_checkpoint'] or resume_preempt - r_file = args['meta']['read_checkpoint'] - copy_data = args['meta']['copy_data'] - pred_depth = args['meta']['pred_depth'] - pred_emb_dim = args['meta']['pred_emb_dim'] + use_bfloat16 = args["meta"]["use_bfloat16"] + model_name = args["meta"]["model_name"] + load_model = args["meta"]["load_checkpoint"] or resume_preempt + r_file = args["meta"]["read_checkpoint"] + copy_data = args["meta"]["copy_data"] + pred_depth = args["meta"]["pred_depth"] + pred_emb_dim = args["meta"]["pred_emb_dim"] if not torch.cuda.is_available(): - device = torch.device('cpu') + device = torch.device("cpu") else: - device = torch.device('cuda:0') + device = torch.device("cuda:0") torch.cuda.set_device(device) # -- DATA - use_gaussian_blur = args['data']['use_gaussian_blur'] - use_horizontal_flip = args['data']['use_horizontal_flip'] - use_color_distortion = args['data']['use_color_distortion'] - color_jitter = args['data']['color_jitter_strength'] + use_gaussian_blur = args["data"]["use_gaussian_blur"] + use_horizontal_flip = args["data"]["use_horizontal_flip"] + use_color_distortion = args["data"]["use_color_distortion"] + color_jitter = args["data"]["color_jitter_strength"] # -- - batch_size = args['data']['batch_size'] - pin_mem = args['data']['pin_mem'] - num_workers = args['data']['num_workers'] - root_path = args['data']['root_path'] - image_folder = args['data']['image_folder'] - crop_size = args['data']['crop_size'] - crop_scale = args['data']['crop_scale'] + batch_size = args["data"]["batch_size"] + pin_mem = args["data"]["pin_mem"] + num_workers = args["data"]["num_workers"] + root_path = args["data"]["root_path"] + image_folder = args["data"]["image_folder"] + crop_size = args["data"]["crop_size"] + crop_scale = args["data"]["crop_scale"] # -- # -- MASK - allow_overlap = args['mask']['allow_overlap'] # whether to allow overlap b/w context and target blocks - patch_size = args['mask']['patch_size'] # patch-size for model training - num_enc_masks = args['mask']['num_enc_masks'] # number of context blocks - min_keep = args['mask']['min_keep'] # min number of patches in context block - enc_mask_scale = args['mask']['enc_mask_scale'] # scale of context blocks - num_pred_masks = args['mask']['num_pred_masks'] # number of target blocks - pred_mask_scale = args['mask']['pred_mask_scale'] # scale of target blocks - aspect_ratio = args['mask']['aspect_ratio'] # aspect ratio of target blocks + allow_overlap = args["mask"][ + "allow_overlap" + ] # whether to allow overlap b/w context and target blocks + patch_size = args["mask"]["patch_size"] # patch-size for model training + num_enc_masks = args["mask"]["num_enc_masks"] # number of context blocks + min_keep = args["mask"]["min_keep"] # min number of patches in context block + enc_mask_scale = args["mask"]["enc_mask_scale"] # scale of context blocks + num_pred_masks = args["mask"]["num_pred_masks"] # number of target blocks + pred_mask_scale = args["mask"]["pred_mask_scale"] # scale of target blocks + aspect_ratio = args["mask"]["aspect_ratio"] # aspect ratio of target blocks # -- # -- OPTIMIZATION - ema = args['optimization']['ema'] - ipe_scale = args['optimization']['ipe_scale'] # scheduler scale factor (def: 1.0) - wd = float(args['optimization']['weight_decay']) - final_wd = float(args['optimization']['final_weight_decay']) - num_epochs = args['optimization']['epochs'] - warmup = args['optimization']['warmup'] - start_lr = args['optimization']['start_lr'] - lr = args['optimization']['lr'] - final_lr = args['optimization']['final_lr'] + ema = args["optimization"]["ema"] + ipe_scale = args["optimization"]["ipe_scale"] # scheduler scale factor (def: 1.0) + wd = float(args["optimization"]["weight_decay"]) + final_wd = float(args["optimization"]["final_weight_decay"]) + num_epochs = args["optimization"]["epochs"] + warmup = args["optimization"]["warmup"] + start_lr = args["optimization"]["start_lr"] + lr = args["optimization"]["lr"] + final_lr = args["optimization"]["final_lr"] # -- LOGGING - folder = args['logging']['folder'] - tag = args['logging']['write_tag'] + folder = args["logging"]["folder"] + tag = args["logging"]["write_tag"] - dump = os.path.join(folder, 'params-ijepa.yaml') - with open(dump, 'w') as f: + dump = os.path.join(folder, "params-ijepa.yaml") + with open(dump, "w") as f: yaml.dump(args, f) # ----------------------------------------------------------------------- # try: - mp.set_start_method('spawn') + mp.set_start_method("spawn") except Exception: pass # -- init torch distributed backend world_size, rank = init_distributed() - logger.info(f'Initialized (rank/world-size) {rank}/{world_size}') + logger.info(f"Initialized (rank/world-size) {rank}/{world_size}") if rank > 0: logger.setLevel(logging.ERROR) # -- log/checkpointing paths - log_file = os.path.join(folder, f'{tag}_r{rank}.csv') - save_path = os.path.join(folder, f'{tag}' + '-ep{epoch}.pth.tar') - latest_path = os.path.join(folder, f'{tag}-latest.pth.tar') + log_file = os.path.join(folder, f"{tag}_r{rank}.csv") + save_path = os.path.join(folder, f"{tag}" + "-ep{epoch}.pth.tar") + latest_path = os.path.join(folder, f"{tag}-latest.pth.tar") load_path = None if load_model: load_path = os.path.join(folder, r_file) if r_file is not None else latest_path # -- make csv_logger - csv_logger = CSVLogger(log_file, - ('%d', 'epoch'), - ('%d', 'itr'), - ('%.5f', 'loss'), - ('%.5f', 'mask-A'), - ('%.5f', 'mask-B'), - ('%d', 'time (ms)')) + csv_logger = CSVLogger( + log_file, + ("%d", "epoch"), + ("%d", "itr"), + ("%.5f", "loss"), + ("%.5f", "mask-A"), + ("%.5f", "mask-B"), + ("%d", "time (ms)"), + ) # -- init model encoder, predictor = init_model( @@ -165,7 +158,8 @@ def main(args, resume_preempt=False): crop_size=crop_size, pred_depth=pred_depth, pred_emb_dim=pred_emb_dim, - model_name=model_name) + model_name=model_name, + ) target_encoder = copy.deepcopy(encoder) # -- make data transforms @@ -178,7 +172,8 @@ def main(args, resume_preempt=False): nenc=num_enc_masks, npred=num_pred_masks, allow_overlap=allow_overlap, - min_keep=min_keep) + min_keep=min_keep, + ) transform = make_transforms( crop_size=crop_size, @@ -186,22 +181,21 @@ def main(args, resume_preempt=False): gaussian_blur=use_gaussian_blur, horizontal_flip=use_horizontal_flip, color_distortion=use_color_distortion, - color_jitter=color_jitter) + color_jitter=color_jitter, + ) # -- init data-loaders/samplers - _, unsupervised_loader, unsupervised_sampler = make_imagenet1k( - transform=transform, - batch_size=batch_size, - collator=mask_collator, - pin_mem=pin_mem, - training=True, - num_workers=num_workers, - world_size=world_size, - rank=rank, - root_path=root_path, - image_folder=image_folder, - copy_data=copy_data, - drop_last=True) + _, unsupervised_loader, unsupervised_sampler = make_iwildcam( + transform=transform, + batch_size=batch_size, + collator=mask_collator, + pin_mem=pin_mem, + num_workers=num_workers, + world_size=world_size, + rank=rank, + root_path=root_path, + drop_last=True, + ) ipe = len(unsupervised_loader) # -- init optimizer and scheduler @@ -217,7 +211,8 @@ def main(args, resume_preempt=False): warmup=warmup, num_epochs=num_epochs, ipe_scale=ipe_scale, - use_bfloat16=use_bfloat16) + use_bfloat16=use_bfloat16, + ) encoder = DistributedDataParallel(encoder, static_graph=True) predictor = DistributedDataParallel(predictor, static_graph=True) target_encoder = DistributedDataParallel(target_encoder) @@ -225,21 +220,26 @@ def main(args, resume_preempt=False): p.requires_grad = False # -- momentum schedule - momentum_scheduler = (ema[0] + i*(ema[1]-ema[0])/(ipe*num_epochs*ipe_scale) - for i in range(int(ipe*num_epochs*ipe_scale)+1)) + momentum_scheduler = ( + ema[0] + i * (ema[1] - ema[0]) / (ipe * num_epochs * ipe_scale) + for i in range(int(ipe * num_epochs * ipe_scale) + 1) + ) start_epoch = 0 # -- load training checkpoint if load_model: - encoder, predictor, target_encoder, optimizer, scaler, start_epoch = load_checkpoint( - device=device, - r_path=load_path, - encoder=encoder, - predictor=predictor, - target_encoder=target_encoder, - opt=optimizer, - scaler=scaler) - for _ in range(start_epoch*ipe): + encoder, predictor, target_encoder, optimizer, scaler, start_epoch = ( + load_checkpoint( + device=device, + r_path=load_path, + encoder=encoder, + predictor=predictor, + target_encoder=target_encoder, + opt=optimizer, + scaler=scaler, + ) + ) + for _ in range(start_epoch * ipe): scheduler.step() wd_scheduler.step() next(momentum_scheduler) @@ -247,25 +247,25 @@ def main(args, resume_preempt=False): def save_checkpoint(epoch): save_dict = { - 'encoder': encoder.state_dict(), - 'predictor': predictor.state_dict(), - 'target_encoder': target_encoder.state_dict(), - 'opt': optimizer.state_dict(), - 'scaler': None if scaler is None else scaler.state_dict(), - 'epoch': epoch, - 'loss': loss_meter.avg, - 'batch_size': batch_size, - 'world_size': world_size, - 'lr': lr + "encoder": encoder.state_dict(), + "predictor": predictor.state_dict(), + "target_encoder": target_encoder.state_dict(), + "opt": optimizer.state_dict(), + "scaler": None if scaler is None else scaler.state_dict(), + "epoch": epoch, + "loss": loss_meter.avg, + "batch_size": batch_size, + "world_size": world_size, + "lr": lr, } if rank == 0: torch.save(save_dict, latest_path) if (epoch + 1) % checkpoint_freq == 0: - torch.save(save_dict, save_path.format(epoch=f'{epoch + 1}')) + torch.save(save_dict, save_path.format(epoch=f"{epoch + 1}")) # -- TRAINING LOOP for epoch in range(start_epoch, num_epochs): - logger.info('Epoch %d' % (epoch + 1)) + logger.info("Epoch %d" % (epoch + 1)) # -- update distributed-data-loader epoch unsupervised_sampler.set_epoch(epoch) @@ -283,6 +283,7 @@ def main(args, resume_preempt=False): masks_1 = [u.to(device, non_blocking=True) for u in masks_enc] masks_2 = [u.to(device, non_blocking=True) for u in masks_pred] return (imgs, masks_1, masks_2) + imgs, masks_enc, masks_pred = load_imgs() maskA_meter.update(len(masks_enc[0][0])) maskB_meter.update(len(masks_pred[0][0])) @@ -313,7 +314,9 @@ def main(args, resume_preempt=False): return loss # Step 1. Forward - with torch.cuda.amp.autocast(dtype=torch.bfloat16, enabled=use_bfloat16): + with torch.cuda.amp.autocast( + dtype=torch.bfloat16, enabled=use_bfloat16 + ): h = forward_target() z = forward_context() loss = loss_fn(z, h) @@ -332,47 +335,62 @@ def main(args, resume_preempt=False): # Step 3. momentum update of target encoder with torch.no_grad(): m = next(momentum_scheduler) - for param_q, param_k in zip(encoder.parameters(), target_encoder.parameters()): - param_k.data.mul_(m).add_((1.-m) * param_q.detach().data) + for param_q, param_k in zip( + encoder.parameters(), target_encoder.parameters() + ): + param_k.data.mul_(m).add_((1.0 - m) * param_q.detach().data) return (float(loss), _new_lr, _new_wd, grad_stats) + (loss, _new_lr, _new_wd, grad_stats), etime = gpu_timer(train_step) loss_meter.update(loss) time_meter.update(etime) # -- Logging def log_stats(): - csv_logger.log(epoch + 1, itr, loss, maskA_meter.val, maskB_meter.val, etime) + csv_logger.log( + epoch + 1, itr, loss, maskA_meter.val, maskB_meter.val, etime + ) if (itr % log_freq == 0) or np.isnan(loss) or np.isinf(loss): - logger.info('[%d, %5d] loss: %.3f ' - 'masks: %.1f %.1f ' - '[wd: %.2e] [lr: %.2e] ' - '[mem: %.2e] ' - '(%.1f ms)' - % (epoch + 1, itr, - loss_meter.avg, - maskA_meter.avg, - maskB_meter.avg, - _new_wd, - _new_lr, - torch.cuda.max_memory_allocated() / 1024.**2, - time_meter.avg)) + logger.info( + "[%d, %5d] loss: %.3f " + "masks: %.1f %.1f " + "[wd: %.2e] [lr: %.2e] " + "[mem: %.2e] " + "(%.1f ms)" + % ( + epoch + 1, + itr, + loss_meter.avg, + maskA_meter.avg, + maskB_meter.avg, + _new_wd, + _new_lr, + torch.cuda.max_memory_allocated() / 1024.0**2, + time_meter.avg, + ) + ) if grad_stats is not None: - logger.info('[%d, %5d] grad_stats: [%.2e %.2e] (%.2e, %.2e)' - % (epoch + 1, itr, - grad_stats.first_layer, - grad_stats.last_layer, - grad_stats.min, - grad_stats.max)) + logger.info( + "[%d, %5d] grad_stats: [%.2e %.2e] (%.2e, %.2e)" + % ( + epoch + 1, + itr, + grad_stats.first_layer, + grad_stats.last_layer, + grad_stats.min, + grad_stats.max, + ) + ) log_stats() - assert not np.isnan(loss), 'loss is nan' + assert not np.isnan(loss), "loss is nan" # -- Save Checkpoint after every epoch - logger.info('avg. loss %.3f' % loss_meter.avg) - save_checkpoint(epoch+1) + logger.info("avg. loss %.3f" % loss_meter.avg) + save_checkpoint(epoch + 1) if __name__ == "__main__":