Fixed conflicts between cluster and github versions
This commit is contained in:
@@ -1,3 +1,5 @@
|
|||||||
*.swp
|
*.swp
|
||||||
*.swo
|
*.swo
|
||||||
__pycache__
|
__pycache__
|
||||||
|
.venv/
|
||||||
|
experiment_logs/
|
||||||
|
|||||||
@@ -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
|
||||||
+29
-27
@@ -21,46 +21,43 @@ logger = logging.getLogger()
|
|||||||
|
|
||||||
|
|
||||||
parser = argparse.ArgumentParser()
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("--folder", type=str, help="location to save submitit logs")
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
'--folder', type=str,
|
"--batch-launch",
|
||||||
help='location to save submitit logs')
|
action="store_true",
|
||||||
|
help="whether fname points to a file to batch-lauch several config files",
|
||||||
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
'--batch-launch', action='store_true',
|
"--fname",
|
||||||
help='whether fname points to a file to batch-lauch several config files')
|
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(
|
parser.add_argument(
|
||||||
'--fname', type=str,
|
"--nodes", type=int, default=1, help="num. nodes to request for job"
|
||||||
help='yaml file containing config file names to launch',
|
)
|
||||||
default='configs.yaml')
|
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
'--partition', type=str,
|
"--tasks-per-node", type=int, default=1, help="num. procs to per node"
|
||||||
help='cluster partition to submit jobs on')
|
)
|
||||||
parser.add_argument(
|
parser.add_argument("--time", type=int, default=4300, help="time in minutes to run job")
|
||||||
'--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:
|
class Trainer:
|
||||||
|
def __init__(self, fname="configs.yaml", load_model=None):
|
||||||
def __init__(self, fname='configs.yaml', load_model=None):
|
|
||||||
self.fname = fname
|
self.fname = fname
|
||||||
self.load_model = load_model
|
self.load_model = load_model
|
||||||
|
|
||||||
def __call__(self):
|
def __call__(self):
|
||||||
fname = self.fname
|
fname = self.fname
|
||||||
load_model = self.load_model
|
load_model = self.load_model
|
||||||
logger.info(f'called-params {fname}')
|
logger.info(f"called-params {fname}")
|
||||||
|
|
||||||
# -- load script params
|
# -- load script params
|
||||||
params = None
|
params = None
|
||||||
with open(fname, 'r') as y_file:
|
with open(fname, "r") as y_file:
|
||||||
params = yaml.load(y_file, Loader=yaml.FullLoader)
|
params = yaml.load(y_file, Loader=yaml.FullLoader)
|
||||||
logger.info('loaded params...')
|
logger.info("loaded params...")
|
||||||
pp = pprint.PrettyPrinter(indent=4)
|
pp = pprint.PrettyPrinter(indent=4)
|
||||||
pp.pprint(params)
|
pp.pprint(params)
|
||||||
|
|
||||||
@@ -69,7 +66,9 @@ class Trainer:
|
|||||||
|
|
||||||
def checkpoint(self):
|
def checkpoint(self):
|
||||||
fb_trainer = Trainer(self.fname, True)
|
fb_trainer = Trainer(self.fname, True)
|
||||||
return submitit.helpers.DelayedSubmission(fb_trainer,)
|
return submitit.helpers.DelayedSubmission(
|
||||||
|
fb_trainer,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def launch():
|
def launch():
|
||||||
@@ -83,7 +82,8 @@ def launch():
|
|||||||
nodes=args.nodes,
|
nodes=args.nodes,
|
||||||
ntasks_per_node=args.tasks_per_node,
|
ntasks_per_node=args.tasks_per_node,
|
||||||
cpus_per_task=10,
|
cpus_per_task=10,
|
||||||
gpus_per_node=args.tasks_per_node)
|
gpus_per_node=args.tasks_per_node,
|
||||||
|
)
|
||||||
|
|
||||||
config_fnames = [args.fname]
|
config_fnames = [args.fname]
|
||||||
|
|
||||||
@@ -91,7 +91,9 @@ def launch():
|
|||||||
with executor.batch():
|
with executor.batch():
|
||||||
for cf in config_fnames:
|
for cf in config_fnames:
|
||||||
fb_trainer = Trainer(cf)
|
fb_trainer = Trainer(cf)
|
||||||
job = executor.submit(fb_trainer,)
|
job = executor.submit(
|
||||||
|
fb_trainer,
|
||||||
|
)
|
||||||
trainers.append(fb_trainer)
|
trainers.append(fb_trainer)
|
||||||
jobs.append(job)
|
jobs.append(job)
|
||||||
|
|
||||||
@@ -99,6 +101,6 @@ def launch():
|
|||||||
print(job.job_id)
|
print(job.job_id)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == '__main__':
|
if __name__ == "__main__":
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
launch()
|
launch()
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -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
|
||||||
@@ -0,0 +1,177 @@
|
|||||||
|
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"]
|
||||||
|
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(
|
||||||
|
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
|
||||||
|
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. 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"]
|
||||||
|
)
|
||||||
|
|
||||||
|
# -- 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. 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()
|
||||||
|
|
||||||
|
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 + 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()
|
||||||
Reference in New Issue
Block a user