Download evaluate.py from ChatterjeeLab/pCoMole: direct link, hf CLI and curl.
- Browser
- Download file 1.26 kB
-
https://huggingface.co/ChatterjeeLab/pCoMole/resolve/main/evaluate.py
- Command line
-
hf download hf://ChatterjeeLab/pCoMole/evaluate.py
-
curl -L -o evaluate.py https://huggingface.co/ChatterjeeLab/pCoMole/resolve/main/evaluate.py
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() | |