""" 评估器模块 — 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: """ 翻译模型评估器。 TODO [Person D]: 实现以下方法。 """ def __init__(self, model: nn.Module, tokenizer, config: dict): """ TODO [Person D]: 1. 保存 model, tokenizer, config 2. 从 config 读取解码策略和评估指标配置 3. 根据策略选择解码函数 """ self.model = model self.tokenizer = tokenizer self.config = 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 evaluate( self, dataloader: DataLoader, src_texts: Optional[list[str]] = None, ref_texts: Optional[list[str]] = None, ) -> dict: """ 在给定数据上进行评估。 TODO [Person D]: 实现以下逻辑: 1. model.eval() 2. 遍历 dataloader,使用选定的解码策略生成翻译 3. 将生成的 token ids 解码为文本 4. 调用 compute_all_metrics 计算指标 5. 返回评估结果 dict """ self.model.eval() device = next(self.model.parameters()).device hypotheses = [] with torch.no_grad(): for batch in tqdm(dataloader, desc="Evaluating"): 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]: """ 翻译一批文本。 TODO [Person D]: 1. tokenize 输入文本 2. 调用解码函数生成翻译 3. 解码为文本 4. 返回翻译结果列表 """ self.model.eval() device = next(self.model.parameters()).device # 1. Tokenize 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) # 2. Decode with torch.no_grad(): output_ids = self.decode_fn(src_ids, src_padding_mask) # 3. Convert to text 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]