#!/usr/bin/env python3 from __future__ import annotations import argparse import json import os import sys import time from collections import defaultdict from pathlib import Path from typing import Any import torch import torch.nn.functional as F def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser( description=( "Evaluate arbitrary continuation-style multiple-choice " "datasets with raw conditional log likelihood." ) ) parser.add_argument("--model-dir", type=Path, default=Path(".")) source = parser.add_mutually_exclusive_group(required=True) source.add_argument("--jsonl", type=Path) source.add_argument("--hf-dataset", type=str) parser.add_argument("--dataset-config", type=str) parser.add_argument("--split", type=str, default="test") parser.add_argument("--context-field", default="ctx") parser.add_argument("--choices-field", default="endings") parser.add_argument("--label-field", default="label") parser.add_argument( "--metadata-field", default="metadata", ) parser.add_argument( "--group-by", nargs="*", default=[], ) parser.add_argument( "--backend", choices=("auto", "torch", "triton"), default="auto", ) parser.add_argument( "--device", choices=("cpu", "cuda"), default="cuda", ) parser.add_argument( "--dtype", choices=("float32", "float16", "bfloat16"), default="bfloat16", ) parser.add_argument("--batch-size", type=int, default=64) parser.add_argument("--max-length", type=int, default=1024) parser.add_argument("--limit", type=int) parser.add_argument("--progress-every", type=int, default=500) parser.add_argument("--output", type=Path, required=True) return parser.parse_args() def load_rows(args: argparse.Namespace) -> list[dict[str, Any]]: if args.jsonl is not None: rows = [] with args.jsonl.open("r", encoding="utf-8") as handle: for line in handle: line = line.strip() if line: rows.append(json.loads(line)) else: from datasets import load_dataset dataset = load_dataset( args.hf_dataset, args.dataset_config, split=args.split, ) rows = [dict(row) for row in dataset] if args.limit is not None: rows = rows[: args.limit] if not rows: raise RuntimeError("The dataset contains no examples") return rows def extract_state_dict(payload: Any) -> dict[str, torch.Tensor]: if isinstance(payload, dict): for container in ("model", "state_dict", "model_state_dict"): candidate = payload.get(container) if isinstance(candidate, dict): payload = candidate break if not isinstance(payload, dict): raise TypeError( f"Unsupported checkpoint type: {type(payload).__name__}" ) result: dict[str, torch.Tensor] = {} for original_name, value in payload.items(): if not torch.is_tensor(value): continue name = str(original_name) for prefix in ("_orig_mod.", "module.", "model."): while name.startswith(prefix): name = name[len(prefix):] result[name] = value return result def encode(tokenizer: Any, text: str) -> list[int]: try: tokens = tokenizer.encode( text, bos=False, eos=False, ) except TypeError: tokens = tokenizer.encode(text) if torch.is_tensor(tokens): tokens = tokens.detach().cpu().reshape(-1).tolist() return [int(token) for token in tokens] def prepare_candidate( tokenizer: Any, context: str, continuation: str, max_length: int, ) -> tuple[list[int], int]: context_tokens = encode(tokenizer, context) continuation_tokens = encode(tokenizer, continuation) if not continuation_tokens: raise ValueError("A continuation encoded to zero tokens") keep_context = max_length - len(continuation_tokens) if keep_context < 1: raise ValueError( "Continuation is longer than the maximum sequence length" ) context_tokens = context_tokens[-keep_context:] tokens = context_tokens + continuation_tokens return tokens, len(continuation_tokens) def main() -> None: args = parse_args() model_dir = args.model_dir.resolve() os.environ["MULTISCREEN_BACKEND"] = args.backend os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") sys.path.insert(0, str(model_dir / "runtime")) from litgpt import Config, GPT, Tokenizer if args.device == "cuda" and not torch.cuda.is_available(): raise RuntimeError("CUDA is unavailable") device = torch.device(args.device) dtype = { "float32": torch.float32, "float16": torch.float16, "bfloat16": torch.bfloat16, }[args.dtype] config = Config.from_file( model_dir / "model_config.yaml" ) model = GPT(config) try: payload = torch.load( model_dir / "lit_model.pth", map_location="cpu", weights_only=False, ) except TypeError: payload = torch.load( model_dir / "lit_model.pth", map_location="cpu", ) incompatible = model.load_state_dict( extract_state_dict(payload), strict=False, ) if incompatible.missing_keys or incompatible.unexpected_keys: raise RuntimeError( "Checkpoint mismatch:\n" f"Missing: {incompatible.missing_keys}\n" f"Unexpected: {incompatible.unexpected_keys}" ) model = model.to(device=device, dtype=dtype) model.eval() tokenizer = Tokenizer(model_dir) rows = load_rows(args) buckets: dict[ int, list[tuple[int, int, list[int], int]], ] = defaultdict(list) labels: list[int] = [] choice_counts: list[int] = [] for example_index, row in enumerate(rows): context = str(row[args.context_field]) choices = list(row[args.choices_field]) label = int(row[args.label_field]) labels.append(label) choice_counts.append(len(choices)) for choice_index, continuation in enumerate(choices): tokens, continuation_length = prepare_candidate( tokenizer, context, str(continuation), min(args.max_length, config.block_size), ) buckets[len(tokens)].append( ( example_index, choice_index, tokens, continuation_length, ) ) scores = [ [float("-inf")] * count for count in choice_counts ] total_candidates = sum( len(records) for records in buckets.values() ) completed = 0 started = time.perf_counter() with torch.inference_mode(): for sequence_length in sorted(buckets): records = buckets[sequence_length] for start in range(0, len(records), args.batch_size): batch = records[start : start + args.batch_size] input_ids = torch.tensor( [record[2] for record in batch], dtype=torch.long, device=device, ) output = model(input_ids) logits = output[0] if isinstance(output, tuple) else output log_probs = F.log_softmax( logits.float(), dim=-1, ) for row_index, record in enumerate(batch): ( example_index, choice_index, tokens, continuation_length, ) = record continuation_start = ( len(tokens) - continuation_length ) target_positions = torch.arange( continuation_start, len(tokens), device=device, ) targets = input_ids[ row_index, target_positions, ] score = log_probs[ row_index, target_positions - 1, targets, ].sum() scores[example_index][choice_index] = float( score.item() ) completed += len(batch) if ( completed == total_candidates or completed % args.progress_every < len(batch) ): elapsed = time.perf_counter() - started print( f"Scored {completed:,}/{total_candidates:,} " f"candidates | " f"{completed / max(elapsed, 1e-9):.1f}/s", flush=True, ) correct = 0 groups: dict[ str, dict[str, list[int]], ] = { field: defaultdict(lambda: [0, 0]) for field in args.group_by } predictions = [] for index, (row, row_scores) in enumerate(zip(rows, scores)): prediction = max( range(len(row_scores)), key=row_scores.__getitem__, ) is_correct = int(prediction == labels[index]) correct += is_correct predictions.append( { "index": index, "prediction": prediction, "label": labels[index], "correct": bool(is_correct), "scores": row_scores, } ) metadata = row.get(args.metadata_field, {}) or {} for field in args.group_by: value = str(metadata.get(field, "unknown")) groups[field][value][0] += is_correct groups[field][value][1] += 1 total = len(rows) accuracy = correct / total group_results = {} for field, values in groups.items(): group_results[field] = {} for value, (group_correct, group_total) in sorted( values.items() ): group_results[field][value] = { "correct": group_correct, "total": group_total, "accuracy": group_correct / group_total, } chance = sum(1.0 / count for count in choice_counts) / total result = { "dataset": args.hf_dataset or str(args.jsonl), "split": args.split, "metric": "raw_continuation_log_likelihood_accuracy", "backend": args.backend, "device": args.device, "dtype": args.dtype, "correct": correct, "total": total, "accuracy": accuracy, "accuracy_percent": accuracy * 100.0, "random_chance": chance, "random_chance_percent": chance * 100.0, "groups": group_results, "predictions": predictions, } args.output.parent.mkdir(parents=True, exist_ok=True) args.output.write_text( json.dumps(result, indent=2) + "\n", encoding="utf-8", ) print() print("=" * 72) print("CUSTOM MULTIPLE-CHOICE RESULTS") print("=" * 72) print( f"Accuracy: {accuracy * 100:.2f}% " f"({correct}/{total})" ) print(f"Random chance: {chance * 100:.2f}%") for field, values in group_results.items(): print(f"\nBy {field}:") for value, metrics in values.items(): print( f" {value}: " f"{metrics['accuracy'] * 100:.2f}% " f"({metrics['correct']}/{metrics['total']})" ) print(f"\nResults: {args.output}") if __name__ == "__main__": main()