""" 评估器模块 — 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. 根据策略选择解码函数 """ raise NotImplementedError("TODO: Person D 实现 Evaluator.__init__") 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 """ raise NotImplementedError("TODO: Person D 实现 evaluate") def translate(self, texts: list[str]) -> list[str]: """ 翻译一批文本。 TODO [Person D]: 1. tokenize 输入文本 2. 调用解码函数生成翻译 3. 解码为文本 4. 返回翻译结果列表 """ raise NotImplementedError("TODO: Person D 实现 translate") def translate_single(self, text: str) -> str: """翻译单条文本。""" return self.translate([text])[0]