File size: 1,261 Bytes
7f316fe | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 | import argparse
import torch
import pytorch_lightning as pl
from paths import add_repo_to_sys_path, resolve_path
from train import build_dataloaders, build_editflow, load_config
add_repo_to_sys_path()
def main():
parser = argparse.ArgumentParser(description="Evaluate an Edit Flow checkpoint on the validation split")
parser.add_argument("--config", type=str, required=True)
parser.add_argument("--ckpt", type=str, required=True)
args = parser.parse_args()
cfg = load_config(args.config)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
editflow, _, _, _, _, _, _ = build_editflow(cfg, device=device)
ckpt = torch.load(resolve_path(args.ckpt), map_location=device, weights_only=False)
state = ckpt["state_dict"] if isinstance(ckpt, dict) and "state_dict" in ckpt else ckpt
editflow.load_state_dict(state, strict=False)
editflow.eval()
_, val_dataloader = build_dataloaders(cfg)
trainer = pl.Trainer(
accelerator="gpu" if torch.cuda.is_available() else "cpu",
devices=1,
logger=False,
enable_checkpointing=False,
)
metrics = trainer.validate(editflow, val_dataloader, verbose=True)
print(metrics)
if __name__ == "__main__":
main()
|