| """ |
| 评估器模块 — Person D 负责实现 |
| |
| 功能要求: |
| 将解码和评估指标整合为统一的评估接口。 |
| |
| 使用方法: |
| evaluator = Evaluator(model, tokenizer, config) |
| results = evaluator.evaluate(test_loader) |
| """ |
|
|
| from __future__ import annotations |
|
|
| import logging |
| from typing import Optional |
|
|
| import torch |
| import torch.nn as nn |
| from torch.utils.data import DataLoader |
| from tqdm import tqdm |
|
|
| from easytranslate.evaluation.metrics import compute_all_metrics |
| from easytranslate.evaluation.decoding import greedy_decode, beam_search_decode, sample_decode |
|
|
| logger = logging.getLogger(__name__) |
|
|
|
|
| class Evaluator: |
| """翻译模型评估器。""" |
|
|
| def __init__(self, model: nn.Module, tokenizer, config: dict): |
| """初始化评估器。""" |
| self.model = model |
| self.tokenizer = tokenizer |
| self.config = config |
|
|
| eval_config = config.get("evaluation", {}) |
| decoding_config = eval_config.get("decoding", {}) |
|
|
| self.strategy = decoding_config.get("strategy", "beam_search") |
| self.max_len = decoding_config.get("max_decode_len", 256) |
| self.beam_size = decoding_config.get("beam_size", 5) |
| self.length_penalty = decoding_config.get("length_penalty", 1.0) |
| self.no_repeat_ngram_size = decoding_config.get("no_repeat_ngram_size", 0) |
|
|
| sampling_config = decoding_config.get("sampling", {}) |
| self.temperature = sampling_config.get("temperature", 1.0) |
| self.top_k = sampling_config.get("top_k", 0) |
| self.top_p = sampling_config.get("top_p", 1.0) |
|
|
| self.metrics = eval_config.get("metrics", ["bleu", "comet", "chrf", "ter"]) |
|
|
| self.bos_id = tokenizer.bos_token_id |
| self.eos_id = tokenizer.eos_token_id |
| self.pad_id = tokenizer.pad_token_id |
|
|
| self.decode_fn = self._get_decode_fn() |
|
|
| def _get_decode_fn(self): |
| """根据策略选择解码函数。""" |
| if self.strategy == "greedy": |
| return self._greedy |
| elif self.strategy == "sampling": |
| return self._sample |
| else: |
| return self._beam_search |
|
|
| def _greedy(self, src_ids, src_padding_mask): |
| return greedy_decode( |
| self.model, src_ids, src_padding_mask, |
| self.bos_id, self.eos_id, max_len=self.max_len, |
| ) |
|
|
| def _beam_search(self, src_ids, src_padding_mask): |
| return beam_search_decode( |
| self.model, src_ids, src_padding_mask, |
| self.bos_id, self.eos_id, |
| beam_size=self.beam_size, max_len=self.max_len, |
| length_penalty=self.length_penalty, |
| no_repeat_ngram_size=self.no_repeat_ngram_size, |
| ) |
|
|
| def _sample(self, src_ids, src_padding_mask): |
| return sample_decode( |
| self.model, src_ids, src_padding_mask, |
| self.bos_id, self.eos_id, |
| max_len=self.max_len, temperature=self.temperature, |
| top_k=self.top_k, top_p=self.top_p, |
| ) |
|
|
| def _validate_batch(self, batch: dict) -> bool: |
| """验证批处理格式是否符合预期。""" |
| required_keys = ["src_ids"] |
| |
| for key in required_keys: |
| if key not in batch: |
| logger.error(f"Batch is missing required key: {key}") |
| return False |
| |
| if not isinstance(batch["src_ids"], torch.Tensor): |
| logger.error(f"src_ids must be a torch.Tensor, got {type(batch['src_ids'])}") |
| return False |
| |
| return True |
|
|
| def evaluate( |
| self, |
| dataloader: DataLoader, |
| src_texts: Optional[list[str]] = None, |
| ref_texts: Optional[list[str]] = None, |
| ) -> dict: |
| """ |
| 在给定数据上进行评估。 |
| |
| Args: |
| dataloader: 测试数据集的 DataLoader |
| src_texts: 源文本列表 (用于某些指标计算) |
| ref_texts: 参考译文列表 |
| |
| Returns: |
| dict: 包含所有评估指标的字典 |
| |
| Raises: |
| ValueError: 如果批处理格式不正确 |
| """ |
| self.model.eval() |
| device = next(self.model.parameters()).device |
|
|
| hypotheses = [] |
|
|
| with torch.no_grad(): |
| for batch in tqdm(dataloader, desc="Evaluating"): |
| |
| if not self._validate_batch(batch): |
| raise ValueError("Invalid batch format. Expected 'src_ids' key with torch.Tensor value.") |
| |
| src_ids = batch["src_ids"].to(device) |
| src_padding_mask = batch.get("src_padding_mask") |
| if src_padding_mask is None: |
| src_padding_mask = src_ids.eq(self.pad_id) |
| else: |
| src_padding_mask = src_padding_mask.to(device) |
|
|
| output_ids = self.decode_fn(src_ids, src_padding_mask) |
|
|
| for i in range(output_ids.size(0)): |
| text = self.tokenizer.decode( |
| output_ids[i].tolist(), skip_special_tokens=True |
| ) |
| hypotheses.append(text) |
|
|
| results = compute_all_metrics( |
| sources=src_texts, |
| hypotheses=hypotheses, |
| references=ref_texts, |
| metrics=self.metrics, |
| ) |
| return results |
|
|
| def translate(self, texts: list[str]) -> list[str]: |
| """翻译一批文本。""" |
| self.model.eval() |
| device = next(self.model.parameters()).device |
|
|
| encoded = [self.tokenizer.encode(t, add_special_tokens=True) for t in texts] |
| max_len_src = max(len(ids) for ids in encoded) |
| src_ids = torch.full((len(texts), max_len_src), self.pad_id, dtype=torch.long, device=device) |
| for i, ids in enumerate(encoded): |
| src_ids[i, :len(ids)] = torch.tensor(ids, dtype=torch.long) |
| src_padding_mask = src_ids.eq(self.pad_id) |
|
|
| with torch.no_grad(): |
| output_ids = self.decode_fn(src_ids, src_padding_mask) |
|
|
| translations = [] |
| for i in range(output_ids.size(0)): |
| text = self.tokenizer.decode(output_ids[i].tolist(), skip_special_tokens=True) |
| translations.append(text) |
|
|
| return translations |
|
|
| def translate_single(self, text: str) -> str: |
| """翻译单条文本。""" |
| return self.translate([text])[0] |