pCoMole / evaluate.py
AlienChen's picture
Upload 83 files
7f316fe verified
Raw History Blame Contribute Delete
1.26 kB
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()