| """ |
| 评估指标模块 — Person D 负责实现 |
| |
| 功能要求: |
| 1. compute_bleu: 计算 SacreBLEU 分数 |
| 2. compute_comet: 计算 COMET 分数 (神经网络评估指标) |
| 3. compute_chrf: 计算 chrF++ 分数 |
| 4. compute_ter: 计算 TER (Translation Edit Rate) |
| 5. compute_all_metrics: 计算所有指标 |
| |
| 技术要点: |
| - SacreBLEU: 标准化的 BLEU 实现,结果可复现 |
| - COMET: 基于预训练语言模型的评估指标,与人类评价相关性最高 |
| - chrF++: 基于字符 n-gram 的 F-score,对中文尤其有用 |
| - TER: 编辑距离,衡量翻译后编辑量 |
| """ |
|
|
| from __future__ import annotations |
|
|
| import logging |
| import time |
| from typing import Optional |
|
|
| logger = logging.getLogger(__name__) |
|
|
|
|
| def _require_sacrebleu(): |
| try: |
| import sacrebleu |
| except ImportError as exc: |
| raise ImportError( |
| "sacrebleu is required for BLEU/chrF/TER evaluation. " |
| "Install it with `pip install sacrebleu`." |
| ) from exc |
| return sacrebleu |
|
|
|
|
| def compute_bleu( |
| hypotheses: list[str], |
| references: list[str], |
| tokenize: str = "zh", |
| ) -> dict: |
| """ |
| 计算 SacreBLEU 分数。 |
| |
| Args: |
| hypotheses: 生成的译文列表 |
| references: 参考译文列表 |
| tokenize: 分词方式,"zh" 表示字符级别分词 |
| |
| Returns: |
| dict: 包含 BLEU 分数及各 n-gram 精度的字典 |
| |
| Raises: |
| ValueError: 如果输入数据为空或长度不匹配 |
| """ |
| if not hypotheses or not references: |
| raise ValueError("hypotheses and references cannot be empty") |
| |
| if len(hypotheses) != len(references): |
| raise ValueError( |
| f"hypotheses and references length mismatch: " |
| f"{len(hypotheses)} vs {len(references)}" |
| ) |
|
|
| sacrebleu = _require_sacrebleu() |
| bleu = sacrebleu.corpus_bleu(hypotheses, [references], tokenize=tokenize) |
| return { |
| "bleu": round(float(bleu.score), 4), |
| "bleu_1": round(float(bleu.precisions[0]), 4), |
| "bleu_2": round(float(bleu.precisions[1]), 4), |
| "bleu_3": round(float(bleu.precisions[2]), 4), |
| "bleu_4": round(float(bleu.precisions[3]), 4), |
| "bp": round(float(bleu.bp), 4), |
| } |
|
|
|
|
| def compute_comet( |
| sources: list[str], |
| hypotheses: list[str], |
| references: list[str], |
| model_name: str = "Unbabel/wmt22-comet-da", |
| batch_size: int = 16, |
| gpus: int = 1, |
| ) -> dict: |
| """ |
| 计算 COMET 分数。 |
| |
| Args: |
| sources: 源文本列表 |
| hypotheses: 生成的译文列表 |
| references: 参考译文列表 |
| model_name: COMET 模型名称 |
| batch_size: 批处理大小 |
| gpus: 使用的 GPU 数量 |
| |
| Returns: |
| dict: 包含 COMET 系统分数和句子级别分数的字典 |
| |
| Raises: |
| ValueError: 如果输入数据为空或长度不匹配 |
| ImportError: 如果 COMET 库未安装 |
| """ |
| try: |
| from comet import download_model, load_from_checkpoint |
| except ImportError: |
| raise ImportError( |
| "COMET is not installed. Please install it with: " |
| "pip install unbabel-comet" |
| ) |
| |
| if not sources or not hypotheses or not references: |
| raise ValueError("sources, hypotheses, and references cannot be empty") |
| |
| if len(sources) != len(hypotheses) or len(hypotheses) != len(references): |
| raise ValueError( |
| f"Input lists have mismatched lengths: " |
| f"sources={len(sources)}, hypotheses={len(hypotheses)}, references={len(references)}" |
| ) |
| |
| model_path = download_model(model_name) |
| model = load_from_checkpoint(model_path) |
|
|
| data = [{"src": s, "mt": h, "ref": r} for s, h, r in zip(sources, hypotheses, references)] |
| prediction = model.predict(data, batch_size=batch_size, gpus=gpus) |
|
|
| system_score = prediction.system_score |
| segment_scores = [float(s) for s in prediction.scores] |
| return {"comet": round(float(system_score), 4), "comet_scores": segment_scores} |
|
|
|
|
| def compute_chrf( |
| hypotheses: list[str], |
| references: list[str], |
| ) -> dict: |
| """ |
| 计算 chrF++ 分数。 |
| |
| Args: |
| hypotheses: 生成的译文列表 |
| references: 参考译文列表 |
| |
| Returns: |
| dict: 包含 chrF++ 分数的字典 |
| |
| Raises: |
| ValueError: 如果输入数据为空或长度不匹配 |
| """ |
| if not hypotheses or not references: |
| raise ValueError("hypotheses and references cannot be empty") |
| |
| if len(hypotheses) != len(references): |
| raise ValueError( |
| f"hypotheses and references length mismatch: " |
| f"{len(hypotheses)} vs {len(references)}" |
| ) |
|
|
| sacrebleu = _require_sacrebleu() |
| chrf = sacrebleu.corpus_chrf(hypotheses, [references]) |
| return {"chrf": round(float(chrf.score), 4)} |
|
|
|
|
| def compute_ter( |
| hypotheses: list[str], |
| references: list[str], |
| ) -> dict: |
| """ |
| 计算 TER 分数。 |
| |
| Args: |
| hypotheses: 生成的译文列表 |
| references: 参考译文列表 |
| |
| Returns: |
| dict: 包含 TER 分数的字典 |
| |
| Raises: |
| ValueError: 如果输入数据为空或长度不匹配 |
| """ |
| if not hypotheses or not references: |
| raise ValueError("hypotheses and references cannot be empty") |
| |
| if len(hypotheses) != len(references): |
| raise ValueError( |
| f"hypotheses and references length mismatch: " |
| f"{len(hypotheses)} vs {len(references)}" |
| ) |
|
|
| sacrebleu = _require_sacrebleu() |
| ter = sacrebleu.corpus_ter(hypotheses, [references]) |
| return {"ter": round(float(ter.score), 4)} |
|
|
|
|
| def compute_all_metrics( |
| sources: Optional[list[str]] = None, |
| hypotheses: Optional[list[str]] = None, |
| references: Optional[list[str]] = None, |
| metrics: list[str] = ["bleu", "comet", "chrf", "ter"], |
| ) -> dict: |
| """ |
| 计算所有指定的评估指标。 |
| |
| Args: |
| sources: 源文本列表 (用于 COMET) |
| hypotheses: 生成的译文列表 |
| references: 参考译文列表 |
| metrics: 需要计算的指标列表 |
| |
| Returns: |
| dict: 包含所有计算指标的字典 |
| |
| Raises: |
| ValueError: 如果 hypotheses 或 references 为空 |
| """ |
| |
| if not hypotheses: |
| raise ValueError("hypotheses cannot be empty") |
| |
| if not references: |
| raise ValueError("references cannot be empty") |
| |
| if len(hypotheses) != len(references): |
| raise ValueError( |
| f"hypotheses and references length mismatch: " |
| f"{len(hypotheses)} vs {len(references)}" |
| ) |
| |
| results = {} |
| |
| for metric in metrics: |
| start = time.time() |
| try: |
| if metric == "bleu": |
| results.update(compute_bleu(hypotheses, references)) |
| elif metric == "comet": |
| if sources is None: |
| logger.warning("COMET requires sources, skipping") |
| continue |
| results.update(compute_comet(sources, hypotheses, references)) |
| elif metric == "chrf": |
| results.update(compute_chrf(hypotheses, references)) |
| elif metric == "ter": |
| results.update(compute_ter(hypotheses, references)) |
| else: |
| logger.warning("Unknown metric: %s, skipping", metric) |
| continue |
| except Exception as e: |
| logger.error(f"Failed to compute {metric}: {str(e)}") |
| continue |
| |
| elapsed = time.time() - start |
| logger.info("Computed %s in %.2fs", metric, elapsed) |
|
|
| return results |