| """ |
| Benchmark Metrics - Evaluation metrics computation |
| """ |
|
|
| import re |
| from typing import List, Dict, Any, Optional |
|
|
|
|
| def extract_number(text: str) -> Optional[float]: |
| """ |
| Extract number from text (for GSM8K and other math problems) |
| |
| Args: |
| text: Input text |
| |
| Returns: |
| Extracted number, or None if not found |
| """ |
| |
| pattern = r"####\s*(-?\d+(?:\.\d+)?)" |
| match = re.search(pattern, text) |
| if match: |
| return float(match.group(1)) |
|
|
| |
| numbers = re.findall(r"-?\d+(?:\.\d+)?", text) |
| if numbers: |
| try: |
| return float(numbers[-1]) |
| except ValueError: |
| pass |
|
|
| return None |
|
|
|
|
| def gsm8k_accuracy( |
| predictions: List[str], |
| ground_truths: List[str], |
| ) -> float: |
| """ |
| Calculate GSM8K accuracy |
| |
| Args: |
| predictions: List of predicted texts |
| ground_truths: List of ground truth answers (including full solution process) |
| |
| Returns: |
| Accuracy (0-1) |
| """ |
| if len(predictions) != len(ground_truths): |
| raise ValueError("Predictions and ground_truths must have the same length") |
|
|
| correct = 0 |
| for pred, gt in zip(predictions, ground_truths): |
| pred_num = extract_number(pred) |
| gt_num = extract_number(gt) |
|
|
| if pred_num is not None and gt_num is not None: |
| if abs(pred_num - gt_num) < 1e-6: |
| correct += 1 |
|
|
| return correct / len(predictions) if predictions else 0.0 |
|
|
|
|
| def humaneval_pass_at_k( |
| results: List[Dict[str, Any]], |
| k: int = 1, |
| ) -> float: |
| """ |
| Calculate HumanEval Pass@k metric |
| |
| Args: |
| results: List of results, each should contain 'output', 'test', 'entry_point' fields |
| k: k value, default 1 |
| |
| Returns: |
| Pass@k score |
| """ |
| |
| |
| |
| return None |
|
|
|
|
| def compute_metrics( |
| outputs: List[Dict[str, Any]], |
| ground_truths: Optional[List[str]] = None, |
| dataset_name: str = "gsm8k", |
| ) -> Dict[str, Any]: |
| """ |
| Compute evaluation metrics |
| |
| Args: |
| outputs: List of generation results |
| ground_truths: List of ground truth answers (optional) |
| dataset_name: Dataset name, used to select appropriate evaluation method |
| |
| Returns: |
| Dictionary of metrics |
| """ |
| metrics = {} |
|
|
| |
| total_tokens = sum(len(o.get("token_ids", [])) for o in outputs) |
| avg_nfe = sum(o.get("nfe", o.get("num_nfes", o.get("n_diff_steps", 0))) for o in outputs) / len(outputs) if outputs else 0 |
| total_time = sum(o.get("generation_time", 0) for o in outputs) |
|
|
| metrics["num_samples"] = len(outputs) |
| metrics["total_tokens"] = total_tokens |
| metrics["avg_tokens_per_sample"] = total_tokens / len(outputs) if outputs else 0 |
| metrics["avg_nfe"] = avg_nfe |
| metrics["total_time"] = total_time |
| metrics["e2e_total_time_s"] = outputs[0].get("e2e_total_time_s", 0.0) if outputs else 0.0 |
| metrics["ttft_s"] = outputs[0].get("ttft_s", 0.0) if outputs else 0.0 |
| metrics["tpot_s"] = outputs[0].get("tpot_s", 0.0) if outputs else 0.0 |
| metrics["e2e_throughput_tok_s"] = outputs[0].get("e2e_throughput_tok_s", 0.0) if outputs else 0.0 |
| metrics["prefill_throughput_tok_s"] = outputs[0].get("prefill_throughput_tok_s", 0.0) if outputs else 0.0 |
| metrics["decode_throughput_tok_s"] = outputs[0].get("decode_throughput_tok_s", 0.0) if outputs else 0.0 |
|
|
| |
| if ground_truths and dataset_name == "gsm8k": |
| predictions = [o.get("text", "") for o in outputs] |
| metrics["accuracy"] = gsm8k_accuracy(predictions, ground_truths) |
| elif ground_truths and dataset_name == "humaneval": |
| |
| metrics["pass_at_1"] = None |
| metrics["note"] = "HumanEval evaluation requires code execution environment" |
|
|
| return metrics |
|
|