add supervised training pipeline

This commit is contained in:
YannAhlgrim
2026-04-15 13:22:44 +02:00
parent e2910cb458
commit 0a38a264f7
3 changed files with 259 additions and 0 deletions
@@ -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
+106
View File
@@ -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()
+124
View File
@@ -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()