""" 评估指标模块 — 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