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