lijn14
完成训练
fb0dadc
Raw
History Blame Contribute Delete
7.65 kB
"""
评估指标模块 — 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