Spaces:
Running
Running
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import math | |
| from pathlib import Path | |
| def choose_device(requested: str): | |
| import torch | |
| if requested == "auto": | |
| return "cuda" if torch.cuda.is_available() else "cpu" | |
| return requested | |
| def main() -> None: | |
| parser = argparse.ArgumentParser(description="Evaluate Ares checkpoint with causal-LM loss/perplexity.") | |
| parser.add_argument("--checkpoint", required=True) | |
| parser.add_argument("--tokenizer", required=True) | |
| parser.add_argument("--eval", nargs="+", required=True) | |
| parser.add_argument("--batch-size", type=int, default=4) | |
| parser.add_argument("--max-batches", type=int, default=50) | |
| parser.add_argument("--device", default="auto") | |
| args = parser.parse_args() | |
| import torch | |
| from torch.utils.data import DataLoader | |
| from .config import AresConfig | |
| from .data import PackedTokenDataset, encode_corpus | |
| from .model import AresForCausalLM | |
| device = choose_device(args.device) | |
| ckpt = torch.load(args.checkpoint, map_location=device) | |
| cfg = AresConfig(**ckpt["config"]) | |
| model = AresForCausalLM(cfg).to(device) | |
| state = {k.replace("_orig_mod.", ""): v for k, v in ckpt["model"].items()} | |
| model.load_state_dict(state, strict=True) | |
| model.eval() | |
| ids = encode_corpus(args.tokenizer, args.eval) | |
| dataset = PackedTokenDataset(ids, cfg.max_seq_len) | |
| loader = DataLoader(dataset, batch_size=args.batch_size, shuffle=False, drop_last=False) | |
| total_loss = 0.0 | |
| total_tokens = 0 | |
| batches = 0 | |
| with torch.no_grad(): | |
| for x, y in loader: | |
| x = x.to(device) | |
| y = y.to(device) | |
| out = model(x, targets=y) | |
| tokens = int(y.numel()) | |
| total_loss += float(out["loss"].detach().cpu()) * tokens | |
| total_tokens += tokens | |
| batches += 1 | |
| if args.max_batches and batches >= args.max_batches: | |
| break | |
| avg_loss = total_loss / max(1, total_tokens) | |
| result = { | |
| "loss": avg_loss, | |
| "perplexity": math.exp(min(20, avg_loss)), | |
| "tokens": total_tokens, | |
| "batches": batches, | |
| } | |
| print(json.dumps(result, indent=2)) | |
| if __name__ == "__main__": | |
| main() | |