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
|
import torch
|
||||||
from logging import getLogger
|
from logging import getLogger
|
||||||
from wilds import get_dataset
|
from wilds import get_dataset
|
||||||
|
from wilds.common.data_loaders import get_eval_loader
|
||||||
|
|
||||||
logger = getLogger()
|
logger = getLogger()
|
||||||
|
|
||||||
@@ -50,6 +51,34 @@ def make_iwildcam(
|
|||||||
return dataset, data_loader, dist_sampler
|
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):
|
class WildsToTorchWrapper(torch.utils.data.Dataset):
|
||||||
"""
|
"""
|
||||||
Mimics the ImageNet wrapper by always returning (image, target).
|
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