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()