evaluation pipeline
This commit is contained in:
@@ -0,0 +1,22 @@
|
||||
meta:
|
||||
seed: 0
|
||||
model_name: vit_huge
|
||||
embed_dim: 1280
|
||||
num_classes: 182
|
||||
patch_size: 14
|
||||
crop_size: 224
|
||||
use_bfloat16: true
|
||||
checkpoint_folder: experiment_logs/supervised-vith14.224-bs.128-ep.300/
|
||||
read_checkpoint: linear_probe-best.pth.tar
|
||||
|
||||
data:
|
||||
batch_size: 128
|
||||
root_path: ./wilds_data
|
||||
num_workers: 10
|
||||
pin_mem: true
|
||||
split: test
|
||||
download: true
|
||||
|
||||
logging:
|
||||
folder: experiment_logs/eval-wilds-vith14/
|
||||
write_tag: iwildcam_test
|
||||
@@ -0,0 +1,102 @@
|
||||
# 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.eval_wilds 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"):
|
||||
self.fname = fname
|
||||
|
||||
def __call__(self):
|
||||
fname = self.fname
|
||||
logger.info(f"called-params {fname}")
|
||||
|
||||
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)
|
||||
|
||||
app_main(args=params)
|
||||
|
||||
def checkpoint(self):
|
||||
fb_trainer = Trainer(self.fname)
|
||||
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()
|
||||
@@ -1,6 +1,7 @@
|
||||
import torch
|
||||
from logging import getLogger
|
||||
from wilds import get_dataset
|
||||
from wilds.common.data_loaders import get_eval_loader
|
||||
|
||||
logger = getLogger()
|
||||
|
||||
@@ -50,6 +51,34 @@ def make_iwildcam(
|
||||
return dataset, data_loader, dist_sampler
|
||||
|
||||
|
||||
def make_iwildcam_eval(
|
||||
transform,
|
||||
batch_size,
|
||||
split="test",
|
||||
num_workers=8,
|
||||
root_path="./wilds_data",
|
||||
download=True,
|
||||
pin_mem=True,
|
||||
):
|
||||
full_dataset = get_dataset(
|
||||
dataset="iwildcam", download=download, root_dir=root_path
|
||||
)
|
||||
|
||||
dataset = full_dataset.get_subset(split, transform=transform)
|
||||
|
||||
logger.info(f"iWildCam {split} eval dataset created with {len(dataset)} samples")
|
||||
|
||||
data_loader = get_eval_loader(
|
||||
"standard",
|
||||
dataset,
|
||||
batch_size=batch_size,
|
||||
num_workers=num_workers,
|
||||
pin_memory=pin_mem,
|
||||
)
|
||||
|
||||
return full_dataset, dataset, data_loader
|
||||
|
||||
|
||||
class WildsToTorchWrapper(torch.utils.data.Dataset):
|
||||
"""
|
||||
Mimics the ImageNet wrapper by always returning (image, target).
|
||||
|
||||
@@ -0,0 +1,182 @@
|
||||
import os
|
||||
import sys
|
||||
import json
|
||||
import yaml
|
||||
import logging
|
||||
|
||||
import numpy as np
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from src.datasets.wilds import make_iwildcam_eval
|
||||
from src.helper import init_model
|
||||
from src.models.head import ViTClassifier
|
||||
from src.transforms import make_transform_eval
|
||||
from src.utils.distributed import init_distributed
|
||||
|
||||
|
||||
_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 strip_module_prefix(state_dict):
|
||||
if not any(k.startswith("module.") for k in state_dict.keys()):
|
||||
return state_dict
|
||||
return {
|
||||
k[len("module.") :] if k.startswith("module.") else k: v
|
||||
for k, v in state_dict.items()
|
||||
}
|
||||
|
||||
|
||||
def _load_yaml(path):
|
||||
with open(path, "r") as f:
|
||||
return yaml.load(f, Loader=yaml.FullLoader)
|
||||
|
||||
|
||||
def _get_seed(args):
|
||||
return int(args.get("meta", {}).get("seed", _GLOBAL_SEED))
|
||||
|
||||
|
||||
def _set_seed(seed):
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
|
||||
|
||||
def _resolve_checkpoint_path(meta_args):
|
||||
if meta_args.get("checkpoint_path"):
|
||||
return meta_args["checkpoint_path"]
|
||||
folder = meta_args.get("checkpoint_folder")
|
||||
fname = meta_args.get("read_checkpoint")
|
||||
if folder is None or fname is None:
|
||||
return None
|
||||
return fname if os.path.isabs(fname) else os.path.join(folder, fname)
|
||||
|
||||
|
||||
def _load_model_state(model, checkpoint_path, device):
|
||||
if checkpoint_path is None:
|
||||
raise ValueError("checkpoint_path or checkpoint_folder/read_checkpoint required")
|
||||
if not os.path.exists(checkpoint_path):
|
||||
raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}")
|
||||
|
||||
checkpoint = torch.load(checkpoint_path, map_location="cpu")
|
||||
if "model" not in checkpoint:
|
||||
raise KeyError(f"No model weights found in checkpoint: {checkpoint_path}")
|
||||
state = strip_module_prefix(checkpoint["model"])
|
||||
msg = model.load_state_dict(state, strict=True)
|
||||
model.to(device)
|
||||
logger.info(f"Loaded model from {checkpoint_path} with msg: {msg}")
|
||||
|
||||
|
||||
def main(args):
|
||||
world_size, rank = init_distributed()
|
||||
|
||||
if dist.is_available() and dist.is_initialized() and world_size > 1 and rank != 0:
|
||||
dist.barrier()
|
||||
dist.destroy_process_group()
|
||||
return
|
||||
|
||||
seed = _get_seed(args)
|
||||
_set_seed(seed)
|
||||
|
||||
meta_args = args.get("meta", {})
|
||||
data_args = args.get("data", {})
|
||||
log_args = args.get("logging", {})
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
device = torch.device("cpu")
|
||||
else:
|
||||
device = torch.device(f"cuda:{torch.cuda.current_device()}")
|
||||
|
||||
folder = log_args.get("folder", "eval_logs")
|
||||
tag = log_args.get("write_tag", "wilds_eval")
|
||||
os.makedirs(folder, exist_ok=True)
|
||||
|
||||
with open(os.path.join(folder, "params-eval.yaml"), "w") as f:
|
||||
yaml.dump(args, f)
|
||||
|
||||
model_name = meta_args["model_name"]
|
||||
embed_dim = meta_args["embed_dim"]
|
||||
num_classes = meta_args["num_classes"]
|
||||
patch_size = meta_args.get("patch_size", 16)
|
||||
crop_size = meta_args.get("crop_size", 224)
|
||||
use_bfloat16 = bool(meta_args.get("use_bfloat16", True))
|
||||
use_autocast = use_bfloat16 and torch.cuda.is_available()
|
||||
|
||||
encoder, _ = init_model(
|
||||
device=device,
|
||||
patch_size=patch_size,
|
||||
crop_size=crop_size,
|
||||
model_name=model_name,
|
||||
)
|
||||
model = ViTClassifier(encoder, num_classes, embed_dim).to(device)
|
||||
|
||||
checkpoint_path = _resolve_checkpoint_path(meta_args)
|
||||
_load_model_state(model, checkpoint_path, device)
|
||||
model.eval()
|
||||
|
||||
eval_transform = make_transform_eval(crop_size=crop_size)
|
||||
split = data_args.get("split", "test")
|
||||
root_path = data_args.get("root_path", "./wilds_data")
|
||||
|
||||
full_dataset, eval_dataset, eval_loader = make_iwildcam_eval(
|
||||
transform=eval_transform,
|
||||
split=split,
|
||||
batch_size=data_args.get("batch_size", 128),
|
||||
root_path=root_path,
|
||||
num_workers=data_args.get("num_workers", 8),
|
||||
pin_mem=bool(data_args.get("pin_mem", True)),
|
||||
download=bool(data_args.get("download", True)),
|
||||
)
|
||||
|
||||
all_y_pred = []
|
||||
all_y_true = []
|
||||
all_metadata = []
|
||||
|
||||
with torch.no_grad():
|
||||
for batch in eval_loader:
|
||||
if len(batch) == 3:
|
||||
imgs, y_true, metadata = batch
|
||||
else:
|
||||
raise ValueError("Expected eval loader to return (x, y, metadata)")
|
||||
imgs = imgs.to(device, non_blocking=True)
|
||||
with torch.cuda.amp.autocast(enabled=use_autocast, dtype=torch.bfloat16):
|
||||
logits = model(imgs)
|
||||
preds = logits.argmax(dim=1).cpu()
|
||||
all_y_pred.append(preds)
|
||||
all_y_true.append(y_true.cpu())
|
||||
all_metadata.append(metadata.cpu())
|
||||
|
||||
if all_y_pred:
|
||||
all_y_pred = torch.cat(all_y_pred, dim=0)
|
||||
all_y_true = torch.cat(all_y_true, dim=0)
|
||||
all_metadata = torch.cat(all_metadata, dim=0) if all_metadata else torch.empty(0)
|
||||
|
||||
if int(all_metadata.shape[0]) != int(all_y_pred.shape[0]):
|
||||
raise ValueError(
|
||||
"Metadata length mismatch with predictions: "
|
||||
f"{int(all_metadata.shape[0])} vs {int(all_y_pred.shape[0])}"
|
||||
)
|
||||
|
||||
metrics = full_dataset.eval(all_y_pred, all_y_true, all_metadata)
|
||||
|
||||
metrics_path = os.path.join(folder, f"{tag}_metrics.json")
|
||||
with open(metrics_path, "w") as f:
|
||||
json.dump(metrics, f, indent=2, sort_keys=True)
|
||||
logger.info(f"Eval metrics saved to {metrics_path}")
|
||||
logger.info(f"Eval metrics: {metrics}")
|
||||
|
||||
if dist.is_available() and dist.is_initialized() and world_size > 1:
|
||||
dist.barrier()
|
||||
dist.destroy_process_group()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise RuntimeError(
|
||||
"Use main_eval_wilds.py to launch this script with a config file."
|
||||
)
|
||||
Reference in New Issue
Block a user