{{ content }}
# Adapted from https://github.com/openai/simple-evals/ """Common utilities for simple evaluations.""" from __future__ import annotations import logging import os import resource import time from collections import defaultdict from dataclasses import dataclass, field from multiprocessing.pool import ThreadPool from typing import Any import httpx import jinja2 import numpy as np import openai import requests from openai import OpenAI from tqdm import tqdm from .constants import MAX_RETRY_ATTEMPTS logger = logging.getLogger(__name__) OPENAI_SYSTEM_MESSAGE_API = "You are a helpful assistant." OPENAI_SYSTEM_MESSAGE_CHATGPT = ( "You are ChatGPT, a large language model trained by OpenAI, based on the GPT-4 architecture." + "\nKnowledge cutoff: 2023-12\nCurrent date: 2024-04-01" ) Message = dict[str, Any] # keys role, content MessageList = list[Message] class SamplerBase: """ Base class for defining a sampling model, which can be evaluated, or used as part of the grading process. """ def __call__(self, message_list: MessageList) -> str: raise NotImplementedError() @dataclass class EvalResult: """Result of running an evaluation (usually consisting of many samples).""" score: float | None # top-line metric metrics: dict[str, float] | None # other metrics htmls: list[str] # strings of valid HTML convos: list[MessageList] # sampled conversations @dataclass class SingleEvalResult: """Result of evaluating a single sample.""" score: float | None metrics: dict[str, float] = field(default_factory=dict) html: str | None = None convo: MessageList | None = None # sampled conversation class Eval: """ Base class for defining an evaluation. """ def __call__(self, sampler: SamplerBase) -> EvalResult: raise NotImplementedError() class LargerHttpxClient(httpx.Client): def __init__(self): timeout_config = httpx.Timeout(3600) limits = httpx.Limits( max_keepalive_connections=3600, max_connections=3600, ) super().__init__(timeout=timeout_config, limits=limits) class ChatCompletionSampler(SamplerBase): """Sample from OpenAI's chat completion API.""" def __init__( self, base_url: str | None = None, model: str | None = None, system_message: str | None = None, temperature: float = 0.0, reasoning_effort: str | None = None, max_tokens: int = 2048, extra_body: dict[str, Any] | None = None, ): self.client = OpenAI(base_url=base_url, http_client=LargerHttpxClient()) if model is None: model = self.client.models.list().data[0].id self.model = model self.system_message = system_message self.temperature = temperature self.max_tokens = max_tokens self.reasoning_effort = reasoning_effort self.extra_body = extra_body self.image_format = "url" logger.debug( "ChatCompletionSampler: model=%s, temp=%.2f, max_tokens=%d", self.model, self.temperature, self.max_tokens, ) def _handle_image( self, image: str, encoding: str = "base64", format: str = "png", ): new_image = { "type": "image_url", "image_url": { "url": f"data:image/{format};{encoding},{image}", }, } return new_image def _handle_text(self, text: str): return {"type": "text", "text": text} def _pack_message(self, role: str, content: Any): return {"role": str(role), "content": content} def __call__(self, message_list: MessageList) -> str: if self.system_message: message_list = [ self._pack_message("system", self.system_message) ] + message_list trial = 0 while trial < MAX_RETRY_ATTEMPTS: try: response = self.client.chat.completions.create( model=self.model, messages=message_list, temperature=self.temperature, max_tokens=self.max_tokens, reasoning_effort=self.reasoning_effort, extra_body=self.extra_body, ) return response.choices[0].message.content or "" except openai.BadRequestError as e: logger.warning("Bad request error: %s", e) return "" except Exception as e: exception_backoff = 2**trial # exponential back off # Log first few retries at debug, later ones at warning log_fn = logger.warning if trial >= 3 else logger.debug log_fn( "Request failed (retry %d/%d, backoff %ds): %s", trial + 1, MAX_RETRY_ATTEMPTS, exception_backoff, e, ) time.sleep(exception_backoff) trial += 1 logger.warning( "All retry attempts exhausted after %d retries, returning empty response", MAX_RETRY_ATTEMPTS, ) return "" QUERY_TEMPLATE_MULTICHOICE = """ Answer the following multiple choice question. The last line of your response should be of the following format: 'Answer: $LETTER' (without quotes) where LETTER is one of ABCD. Think step by step before answering. {Question} A) {A} B) {B} C) {C} D) {D} """.strip() ANSWER_PATTERN_MULTICHOICE = r"(?i)Answer\s*:\s*([A-D])" ANSWER_PATTERN = r"(?i)Answer\s*:\s*([^\n]+)" EQUALITY_TEMPLATE = r""" Look at the following two expressions (answers to a math problem) and judge whether they are equivalent. Only perform trivial simplifications Examples: Expression 1: $2x+3$ Expression 2: $3+2x$ Yes Expression 1: 3/2 Expression 2: 1.5 Yes Expression 1: $x^2+2x+1$ Expression 2: $y^2+2y+1$ No Expression 1: $x^2+2x+1$ Expression 2: $(x+1)^2$ Yes Expression 1: 3245/5 Expression 2: 649 No (these are actually equal, don't mark them equivalent if you need to do nontrivial simplifications) Expression 1: 2/(-3) Expression 2: -2/3 Yes (trivial simplifications are allowed) Expression 1: 72 degrees Expression 2: 72 Yes (give benefit of the doubt to units) Expression 1: 64 Expression 2: 64 square feet Yes (give benefit of the doubt to units) --- YOUR TASK Respond with only "Yes" or "No" (without quotes). Do not include a rationale. Expression 1: %(expression1)s Expression 2: %(expression2)s """.strip() HTML_JINJA = """
Correct Answer: {{ correct_answer }}
Extracted Answer: {{ extracted_answer }}
Score: {{ score }}
""" def format_multichoice_question(row): return QUERY_TEMPLATE_MULTICHOICE.format(**row) def check_equality(sampler: SamplerBase, expr1: str, expr2: str): prompt = EQUALITY_TEMPLATE % {"expression1": expr1, "expression2": expr2} response = sampler([dict(content=prompt, role="user")]) return (response or "").lower().strip() == "yes" def _compute_stat(values: list, stat: str): if stat == "mean": return np.mean(values) elif stat == "std": return np.std(values) elif stat == "min": return np.min(values) elif stat == "max": return np.max(values) else: raise ValueError(f"Unknown {stat =}") def aggregate_results( single_eval_results: list[SingleEvalResult], default_stats: tuple[str, ...] = ("mean", "std"), name2stats: dict[str, tuple[str, ...]] | None = None, ) -> EvalResult: """ Aggregate results from multiple evaluations into a single EvalResult. """ name2stats = name2stats or {} name2values = defaultdict(list) htmls = [] convos = [] for single_eval_result in single_eval_results: # Skip None results if single_eval_result is None: continue for name, value in single_eval_result.metrics.items(): name2values[name].append(value) if single_eval_result.score is not None: name2values["score"].append(single_eval_result.score) htmls.append(single_eval_result.html) convos.append(single_eval_result.convo) final_metrics = {} for name, values in name2values.items(): stats = name2stats.get(name, default_stats) for stat in stats: key = name if stat == "mean" else f"{name}:{stat}" final_metrics[key] = _compute_stat(values, stat) return EvalResult( score=final_metrics.pop("score", None), metrics=final_metrics, htmls=htmls, convos=convos, ) def map_with_progress(f: callable, xs: list[Any], num_threads: int) -> list[Any]: """Apply f to each element of xs, using a ThreadPool, and show progress.""" # Use quiet progress bar that doesn't pollute logs if os.getenv("debug"): return list(map(f, tqdm(xs, total=len(xs), leave=False))) else: with ThreadPool(min(num_threads, len(xs))) as pool: return list(tqdm(pool.imap(f, xs), total=len(xs), leave=False)) jinja_env = jinja2.Environment( loader=jinja2.BaseLoader(), undefined=jinja2.StrictUndefined, autoescape=jinja2.select_autoescape(["html", "xml"]), ) _message_template = """ """ def message_to_html(message: Message) -> str: """ Generate HTML snippet (inside a| Metric | Value |
|---|---|
| Score | {{ score | float | round(3) }} |
| {{ name }} | {{ value }} |