From 8956531fe5df1f6fceeedb40b2243cc472ab395a Mon Sep 17 00:00:00 2001 From: YannAhlgrim Date: Sun, 5 Apr 2026 18:44:54 +0200 Subject: [PATCH 1/6] change explicitly for slurm cluster --- main_distributed.py | 70 +++++++++++++++++++++++---------------------- 1 file changed, 36 insertions(+), 34 deletions(-) diff --git a/main_distributed.py b/main_distributed.py index 7cb3846..9941c73 100644 --- a/main_distributed.py +++ b/main_distributed.py @@ -21,46 +21,43 @@ logger = logging.getLogger() parser = argparse.ArgumentParser() +parser.add_argument("--folder", type=str, help="location to save submitit logs") parser.add_argument( - '--folder', type=str, - help='location to save submitit logs') + "--batch-launch", + action="store_true", + help="whether fname points to a file to batch-lauch several config files", +) parser.add_argument( - '--batch-launch', action='store_true', - help='whether fname points to a file to batch-lauch several config files') + "--fname", + type=str, + help="yaml file containing config file names to launch", + default="configs.yaml", +) +parser.add_argument("--partition", type=str, help="cluster partition to submit jobs on") parser.add_argument( - '--fname', type=str, - help='yaml file containing config file names to launch', - default='configs.yaml') + "--nodes", type=int, default=1, help="num. nodes to request for job" +) parser.add_argument( - '--partition', type=str, - help='cluster partition to submit jobs on') -parser.add_argument( - '--nodes', type=int, default=1, - help='num. nodes to request for job') -parser.add_argument( - '--tasks-per-node', type=int, default=1, - help='num. procs to per node') -parser.add_argument( - '--time', type=int, default=4300, - help='time in minutes to run job') + "--tasks-per-node", type=int, default=1, help="num. procs to per node" +) +parser.add_argument("--time", type=int, default=4300, help="time in minutes to run job") class Trainer: - - def __init__(self, fname='configs.yaml', load_model=None): + def __init__(self, fname="configs.yaml", load_model=None): self.fname = fname self.load_model = load_model def __call__(self): fname = self.fname load_model = self.load_model - logger.info(f'called-params {fname}') + logger.info(f"called-params {fname}") # -- load script params params = None - with open(fname, 'r') as y_file: + with open(fname, "r") as y_file: params = yaml.load(y_file, Loader=yaml.FullLoader) - logger.info('loaded params...') + logger.info("loaded params...") pp = pprint.PrettyPrinter(indent=4) pp.pprint(params) @@ -69,21 +66,24 @@ class Trainer: def checkpoint(self): fb_trainer = Trainer(self.fname, True) - return submitit.helpers.DelayedSubmission(fb_trainer,) + return submitit.helpers.DelayedSubmission( + fb_trainer, + ) def launch(): - executor = submitit.AutoExecutor( - folder=os.path.join(args.folder, 'job_%j'), - slurm_max_num_timeout=20) + executor = submitit.SlurmExecutor( + folder=os.path.join(args.folder, "job_%j"), max_num_timeout=20 + ) executor.update_parameters( - slurm_partition=args.partition, - slurm_mem_per_gpu='55G', - timeout_min=args.time, + partition=args.partition, + mem_per_gpu="55G", + time=args.time, nodes=args.nodes, - tasks_per_node=args.tasks_per_node, + ntasks_per_node=args.tasks_per_node, cpus_per_task=10, - gpus_per_node=args.tasks_per_node) + gpus_per_node=args.tasks_per_node, + ) config_fnames = [args.fname] @@ -91,7 +91,9 @@ def launch(): with executor.batch(): for cf in config_fnames: fb_trainer = Trainer(cf) - job = executor.submit(fb_trainer,) + job = executor.submit( + fb_trainer, + ) trainers.append(fb_trainer) jobs.append(job) @@ -99,6 +101,6 @@ def launch(): print(job.job_id) -if __name__ == '__main__': +if __name__ == "__main__": args = parser.parse_args() launch() From 4ec019f650bafb7a3c93cc64fc37dbd25aa4ee9c Mon Sep 17 00:00:00 2001 From: YannAhlgrim Date: Sun, 5 Apr 2026 18:46:55 +0200 Subject: [PATCH 2/6] submitit for main_distributed.py --- requirements.txt | 2 ++ 1 file changed, 2 insertions(+) diff --git a/requirements.txt b/requirements.txt index 77506e5..c27dafc 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,6 @@ certifi==2026.2.25 charset-normalizer==3.4.6 +cloudpickle==3.1.2 cmake==4.3.1 filelock==3.19.1 idna==3.11 @@ -34,6 +35,7 @@ requests==2.32.5 scikit-learn==1.6.1 scipy==1.13.1 six==1.17.0 +submitit==1.5.4 sympy==1.14.0 threadpoolctl==3.6.0 torch==2.0.1 From e2910cb4589cab3db6b94d44576262b021a1d885 Mon Sep 17 00:00:00 2001 From: YannAhlgrim Date: Sun, 12 Apr 2026 11:19:06 +0200 Subject: [PATCH 3/6] Header for Linear Probing --- src/models/head.py | 22 ++++++++++++++++++++++ 1 file changed, 22 insertions(+) create mode 100644 src/models/head.py diff --git a/src/models/head.py b/src/models/head.py new file mode 100644 index 0000000..f846486 --- /dev/null +++ b/src/models/head.py @@ -0,0 +1,22 @@ +import torch +import torch.nn as nn + + +class ViTClassifier(nn.Module): + def __init__(self, encoder, num_classes, embed_dim): + super().__init__() + self.encoder = encoder + self.head = nn.Linear(embed_dim, num_classes) + + nn.init.trunc_normal_(self.head.weight, std=0.01) + nn.init.zeros_(self.head.bias) + + def forward(self, x): + # ViT -> (B, N, D) + features = self.encoder(x) + + # Average Pool -> (B, D) + avg_embed = features.mean(dim=1) + + logits = self.head(avg_embed) + return logits From 0a38a264f76aed9d257fcd6d7c997a63502ff9a1 Mon Sep 17 00:00:00 2001 From: YannAhlgrim Date: Wed, 15 Apr 2026 13:22:44 +0200 Subject: [PATCH 4/6] add supervised training pipeline --- configs/supervised_wilds_vith14_ep300.yaml | 29 +++++ main_distributed_supervised.py | 106 ++++++++++++++++++ src/train_supervised.py | 124 +++++++++++++++++++++ 3 files changed, 259 insertions(+) create mode 100644 configs/supervised_wilds_vith14_ep300.yaml create mode 100644 main_distributed_supervised.py create mode 100644 src/train_supervised.py diff --git a/configs/supervised_wilds_vith14_ep300.yaml b/configs/supervised_wilds_vith14_ep300.yaml new file mode 100644 index 0000000..53c0adc --- /dev/null +++ b/configs/supervised_wilds_vith14_ep300.yaml @@ -0,0 +1,29 @@ +meta: + model_name: vit_huge + embed_dim: 1280 + load_checkpoint: true + read_checkpoint: jepa-latest.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 + +optimization: + optimizer: adamw # 'adamw' or 'sgd' + freeze_weights: true # true for linear probing, false for full fine-tuning + epochs: 50 + lr: 0.001 + weight_decay: 0.05 + start_lr: 0.0001 + final_lr: 1.0e-06 + warmup: 5 + ipe_scale: 1.0 + +logging: + folder: experiment_logs/supervised-vith14.224-bs.128-ep.300/ + write_tag: linear_probe diff --git a/main_distributed_supervised.py b/main_distributed_supervised.py new file mode 100644 index 0000000..37f2c5e --- /dev/null +++ b/main_distributed_supervised.py @@ -0,0 +1,106 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. +# + +import argparse +import logging +import os +import pprint +import sys +import yaml + +import submitit + +from src.train_supervised import main as app_main + +logging.basicConfig(stream=sys.stdout, level=logging.INFO) +logger = logging.getLogger() + + +parser = argparse.ArgumentParser() +parser.add_argument("--folder", type=str, help="location to save submitit logs") +parser.add_argument( + "--batch-launch", + action="store_true", + help="whether fname points to a file to batch-lauch several config files", +) +parser.add_argument( + "--fname", + type=str, + help="yaml file containing config file names to launch", + default="configs.yaml", +) +parser.add_argument("--partition", type=str, help="cluster partition to submit jobs on") +parser.add_argument( + "--nodes", type=int, default=1, help="num. nodes to request for job" +) +parser.add_argument( + "--tasks-per-node", type=int, default=1, help="num. procs to per node" +) +parser.add_argument("--time", type=int, default=4300, help="time in minutes to run job") + + +class Trainer: + def __init__(self, fname="configs.yaml", load_model=None): + self.fname = fname + self.load_model = load_model + + def __call__(self): + fname = self.fname + load_model = self.load_model + logger.info(f"called-params {fname}") + + # -- load script params + params = None + with open(fname, "r") as y_file: + params = yaml.load(y_file, Loader=yaml.FullLoader) + logger.info("loaded params...") + pp = pprint.PrettyPrinter(indent=4) + pp.pprint(params) + + resume_preempt = False if load_model is None else load_model + app_main(args=params, resume_preempt=resume_preempt) + + def checkpoint(self): + fb_trainer = Trainer(self.fname, True) + return submitit.helpers.DelayedSubmission( + fb_trainer, + ) + + +def launch(): + executor = submitit.SlurmExecutor( + folder=os.path.join(args.folder, "job_%j"), max_num_timeout=20 + ) + executor.update_parameters( + partition=args.partition, + mem_per_gpu="55G", + time=args.time, + nodes=args.nodes, + ntasks_per_node=args.tasks_per_node, + cpus_per_task=10, + gpus_per_node=args.tasks_per_node, + ) + + config_fnames = [args.fname] + + jobs, trainers = [], [] + with executor.batch(): + for cf in config_fnames: + fb_trainer = Trainer(cf) + job = executor.submit( + fb_trainer, + ) + trainers.append(fb_trainer) + jobs.append(job) + + for job in jobs: + print(job.job_id) + + +if __name__ == "__main__": + args = parser.parse_args() + launch() diff --git a/src/train_supervised.py b/src/train_supervised.py new file mode 100644 index 0000000..383971a --- /dev/null +++ b/src/train_supervised.py @@ -0,0 +1,124 @@ +import os +import sys +import yaml +import logging +import torch +import torch.nn as nn +from torch.nn.parallel import DistributedDataParallel +import torch.nn.functional as F +import numpy as np + +from src.datasets.wilds import make_iwildcam +from src.helper import load_checkpoint, init_model, init_opt +from src.transforms import make_transforms +from src.models.head import ViTClassifier # Import our new class +from src.utils.distributed import init_distributed +from src.utils.logging import CSVLogger, AverageMeter + +# -- +log_timings = True +log_freq = 10 +checkpoint_freq = 50 +# -- + +_GLOBAL_SEED = 0 +np.random.seed(_GLOBAL_SEED) +torch.manual_seed(_GLOBAL_SEED) +torch.backends.cudnn.benchmark = True + +logging.basicConfig(stream=sys.stdout, level=logging.INFO) +logger = logging.getLogger() + + +def main(args): + # -- Init Distributed + world_size, rank = init_distributed() + device = torch.device(f"cuda:{torch.cuda.current_device()}") + + # -- Extract Config Params + m_args = args["meta"] + o_args = args["optimization"] + d_args = args["data"] + + # -- 1. Initialize Encoder + encoder, _ = init_model( + device=device, + patch_size=args.get("mask", {}).get("patch_size", 14), + crop_size=d_args["crop_size"], + model_name=m_args["model_name"], + ) + + # -- 2. Wrap in Classification Head + embed_dim = m_args.get("embed_dim") + model = ViTClassifier(encoder, m_args["num_classes"], embed_dim).to(device) + + # -- 3. Handle Freezing (Linear Probing vs Fine-Tuning) + if o_args["freeze_weights"]: + logger.info("Freezing encoder weights (Linear Probing mode)") + for name, param in model.encoder.named_parameters(): + param.requires_grad = False + else: + logger.info("Training full model (Fine-tuning mode)") + + # -- 4. Load Pre-trained Weights + if m_args["load_checkpoint"]: + load_path = os.path.join(args["logging"]["folder"], m_args["read_checkpoint"]) + checkpoint = torch.load(load_path, map_location="cpu") + msg = model.encoder.load_state_dict(checkpoint["encoder"], strict=False) + logger.info(f"Loaded encoder from {load_path} with msg: {msg}") + + # -- 5. Data Setup + transform = make_transforms(crop_size=d_args["crop_size"]) + _, loader, sampler = make_iwildcam( + transform=transform, + split="train", + batch_size=d_args["batch_size"], + root_path=d_args["root_path"], + rank=rank, + world_size=world_size, + collator=None, # No mask collator needed + ) + + # -- 6. Optimizer Selection + params = [p for p in model.parameters() if p.requires_grad] + if o_args["optimizer"].lower() == "adamw": + optimizer = torch.optim.AdamW( + params, lr=o_args["lr"], weight_decay=o_args["weight_decay"] + ) + else: + optimizer = torch.optim.SGD( + params, lr=o_args["lr"], momentum=0.9, weight_decay=o_args["weight_decay"] + ) + + criterion = nn.CrossEntropyLoss().to(device) + model = DistributedDataParallel(model, device_ids=[torch.cuda.current_device()]) + + # -- 7. Training Loop + for epoch in range(o_args["epochs"]): + sampler.set_epoch(epoch) + model.train() + loss_meter = AverageMeter() + + for itr, (imgs, labels, _) in enumerate(loader): + imgs, labels = imgs.to(device), labels.to(device) + + with torch.cuda.amp.autocast( + enabled=m_args["use_bfloat16"], dtype=torch.bfloat16 + ): + outputs = model(imgs) + loss = criterion(outputs, labels) + + optimizer.zero_grad() + loss.backward() + optimizer.step() + + loss_meter.update(loss.item()) + + if itr % 10 == 0 and rank == 0: + logger.info( + f"Epoch {epoch} [{itr}/{len(loader)}] Loss: {loss_meter.avg:.4f}" + ) + + +if __name__ == "__main__": + main() From ebad8ee627c3d3cd14c0f3d2c89e844080ce574e Mon Sep 17 00:00:00 2001 From: YannAhlgrim Date: Wed, 15 Apr 2026 13:23:24 +0200 Subject: [PATCH 5/6] experiment logs and venv --- .gitignore | 2 ++ 1 file changed, 2 insertions(+) diff --git a/.gitignore b/.gitignore index 6d669bf..bfe3776 100644 --- a/.gitignore +++ b/.gitignore @@ -1,3 +1,5 @@ *.swp *.swo __pycache__ +.venv/ +experiment_logs/ From 4c4d6713e06203d36fb0ab1109dfefd35024535d Mon Sep 17 00:00:00 2001 From: YannAhlgrim Date: Wed, 15 Apr 2026 13:43:05 +0200 Subject: [PATCH 6/6] add save checkpoint logic --- src/train_supervised.py | 101 ++++++++++++++++++++++++++++++---------- 1 file changed, 77 insertions(+), 24 deletions(-) diff --git a/src/train_supervised.py b/src/train_supervised.py index 383971a..8821668 100644 --- a/src/train_supervised.py +++ b/src/train_supervised.py @@ -39,6 +39,16 @@ def main(args): m_args = args["meta"] o_args = args["optimization"] d_args = args["data"] + l_args = args["logging"] + + # -- Paths for saving + folder = l_args["folder"] + tag = l_args["write_tag"] + if not os.path.exists(folder): + os.makedirs(folder, exist_ok=True) + + save_path = os.path.join(folder, f"{tag}" + "-ep{epoch}.pth.tar") + latest_path = os.path.join(folder, f"{tag}-latest.pth.tar") # -- 1. Initialize Encoder encoder, _ = init_model( @@ -52,7 +62,7 @@ def main(args): embed_dim = m_args.get("embed_dim") model = ViTClassifier(encoder, m_args["num_classes"], embed_dim).to(device) - # -- 3. Handle Freezing (Linear Probing vs Fine-Tuning) + # -- 3. Handle Freezing if o_args["freeze_weights"]: logger.info("Freezing encoder weights (Linear Probing mode)") for name, param in model.encoder.named_parameters(): @@ -60,26 +70,7 @@ def main(args): else: logger.info("Training full model (Fine-tuning mode)") - # -- 4. Load Pre-trained Weights - if m_args["load_checkpoint"]: - load_path = os.path.join(args["logging"]["folder"], m_args["read_checkpoint"]) - checkpoint = torch.load(load_path, map_location="cpu") - msg = model.encoder.load_state_dict(checkpoint["encoder"], strict=False) - logger.info(f"Loaded encoder from {load_path} with msg: {msg}") - - # -- 5. Data Setup - transform = make_transforms(crop_size=d_args["crop_size"]) - _, loader, sampler = make_iwildcam( - transform=transform, - split="train", - batch_size=d_args["batch_size"], - root_path=d_args["root_path"], - rank=rank, - world_size=world_size, - collator=None, # No mask collator needed - ) - - # -- 6. Optimizer Selection + # -- 4. Optimizer Selection params = [p for p in model.parameters() if p.requires_grad] if o_args["optimizer"].lower() == "adamw": optimizer = torch.optim.AdamW( @@ -90,11 +81,70 @@ def main(args): params, lr=o_args["lr"], momentum=0.9, weight_decay=o_args["weight_decay"] ) + # -- 5. Resume/Load Logic + start_epoch = 0 + # Priority 1: Check if we are resuming from an interrupted run (latest-path) + # Priority 2: Check if we are loading a specific pre-trained checkpoint (m_args["load_checkpoint"]) + + checkpoint_to_load = None + resuming_interrupted = False + + if os.path.exists(latest_path): + checkpoint_to_load = latest_path + resuming_interrupted = True + elif m_args["load_checkpoint"]: + checkpoint_to_load = os.path.join(folder, m_args["read_checkpoint"]) + + if checkpoint_to_load: + checkpoint = torch.load(checkpoint_to_load, map_location="cpu") + + if resuming_interrupted: + # Load full state to resume exactly where we left off + model.load_state_dict(checkpoint["model"]) + optimizer.load_state_dict(checkpoint["opt"]) + start_epoch = checkpoint["epoch"] + logger.info( + f"Resuming training from {checkpoint_to_load} at epoch {start_epoch}" + ) + else: + # Loading just encoder weights for a fresh supervised run + msg = model.encoder.load_state_dict(checkpoint["encoder"], strict=False) + logger.info( + f"Loaded pre-trained encoder from {checkpoint_to_load} with msg: {msg}" + ) + + # -- 6. Data Setup + transform = make_transforms(crop_size=d_args["crop_size"]) + _, loader, sampler = make_iwildcam( + transform=transform, + split="train", + batch_size=d_args["batch_size"], + root_path=d_args["root_path"], + rank=rank, + world_size=world_size, + collator=None, + ) + criterion = nn.CrossEntropyLoss().to(device) model = DistributedDataParallel(model, device_ids=[torch.cuda.current_device()]) - # -- 7. Training Loop - for epoch in range(o_args["epochs"]): + # -- 7. Define Save Function + def save_checkpoint(epoch, current_loss): + save_dict = { + "model": model.module.state_dict(), # model.module because of DDP + "opt": optimizer.state_dict(), + "epoch": epoch, + "loss": current_loss, + "args": args, + } + if rank == 0: + torch.save(save_dict, latest_path) + if epoch % checkpoint_freq == 0: + torch.save(save_dict, save_path.format(epoch=epoch)) + logger.info(f"Checkpoint saved at epoch {epoch}") + + # -- 8. Training Loop + for epoch in range(start_epoch, o_args["epochs"]): sampler.set_epoch(epoch) model.train() loss_meter = AverageMeter() @@ -116,9 +166,12 @@ def main(args): if itr % 10 == 0 and rank == 0: logger.info( - f"Epoch {epoch} [{itr}/{len(loader)}] Loss: {loss_meter.avg:.4f}" + f"Epoch {epoch + 1} [{itr}/{len(loader)}] Loss: {loss_meter.avg:.4f}" ) + # Save at the end of every epoch + save_checkpoint(epoch + 1, loss_meter.avg) + if __name__ == "__main__": main()