integrate make_iwildcam in train.py

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